# 6.2. 矩阵乘法
每一次 Dense 层的前向传播、每一个循环时间步,都归结为同一件事:一次矩阵乘积。线性模型的预测,以及 KNN 和 t-SNE 内部的成对投影,也是如此。RustyML 不把它们交给 ndarray 的 `.dot()`。
RustyML 把它们委托给 [`gemmkit`](https://crates.io/crates/gemmkit),一个纯 Rust 的 GEMM 引擎。crate 经由它那层零拷贝的 [`gemmkit-ndarray`](https://crates.io/crates/gemmkit-ndarray) 适配器够到 gemmkit。`math` feature 只点名适配器,引擎由适配器带进来。RustyML 自己只在 `src/math/matmul.rs` 里留下薄薄一层代码。
本页讲 4 件事。它讲这套后端做了什么。它讲为什么 RustyML 用它而不是 `.dot()`。它讲后端如何在串行与并行之间抉择,并行时又开多宽。它讲唯一一处你能改动的地方:`rustyml::tuning::matmul` 下的运行期调优接口。
crate 自己的 matmul 入口是 `pub(crate)` 的。你没法在自己的代码里调用 `dot_par`。各个估计器是直接够到 `gemmkit_ndarray`。它们并不经过 RustyML 重新导出的任何类型或函数。
你能理解它的行为。它决定你的模型跑多快。你也能针对你的机器重新调整那些阈值。如果你要在自己的代码里做矩阵乘积,用 ndarray 的 `.dot()`(见 [1.3. 使用ndarray准备数据](../Chapter-01/1.3._使用ndarray准备数据.md))。不要直接用这套后端。
## 6.2.1. 这套后端是什么,为什么它是内部的
这里有 2 层,分清楚会省很多事。引擎是 gemmkit。它在带步长的视图上计算 `C <- alpha*A*B + beta*C`。它在运行期挑选自己的指令集。它自己做打包和分块,并且掌握全部调度决策。
`gemmkit-ndarray` 是一层很薄的适配器。它直接从 `ArrayBase<S, Ix2>` 里读出数据指针和步长,转交给引擎。对 C 序视图、F 序视图、一般步长视图,还是负步长视图,它都不做任何拷贝。
这个适配器本身就已经是合适的调用侧 API。所以 RustyML 的层和估计器就直接调它。它们用 `gemmkit_ndarray::dot` 走后端自动调度的分配式乘积。它们用 `gemmkit_ndarray::gemm`,让调用方自己持有输出缓冲区。它们用 `gemmkit_ndarray::gemm_fused`,让偏置和激活搭同一趟顺风车。
`src/math/matmul.rs` 里剩下的,用它自己的话说,就是"crate 对 gemmkit 后端的少数几处补充"。总共只有 4 个条目:
| 条目 | 可见性 | 是什么 |
|---|---|---|
| `dot_par(a, b, par)` | `pub(crate)` | 带显式 `gemmkit_ndarray::Parallelism` 的分配式 `A @ B`(普通的 `dot` 一律用自动默认值) |
| `matvec(a, x, par)` | `pub(crate)` | 操作数为 `Array1` 的 matvec,把 `x` 包成 `[k, 1]` 的列,gemmkit 据此改走它的 GEMV 路径 |
| `gemm_chunk_rows(row_len)` | `pub`,`#[doc(hidden)]` | `gemm_chunk_elems() / row_len`,钳制在 `[16, 4096]` 行之内 |
| `cache_resident::<T>(rows, cols)` | `pub`,`#[doc(hidden)]` | `rows * cols * size_of::<T>()` 是否低于 `cache_resident_max_bytes()` |
前 2 个条目对 `T: gemmkit_ndarray::GemmScalar` 泛化。在 RustyML 的构建里,这恰好就是 `f32` 和 `f64`。gemmkit 另外还能在可选的 `half` feature 下支持 `f16` 和 `bf16`,在 `int8` feature 下支持 `i8`。RustyML 两个都没打开,所以这里既没有半精度,也没有整数 matmul。
RustyML 唯一打开的非默认 feature 是 `epilogue`,开在 `gemmkit-ndarray` 上,为的是那条融合路径。如果你需要 `f16`、`bf16` 或 `i8` 支持,就得自己直接对着 gemmkit 写代码。
后 2 个条目根本不是乘积。它们是调用侧的分块策略,供那些成对投影大到一次装不下、否则就得整块物化的估计器使用。它们在技术上可以通过 `rustyml::math::matmul::gemm_chunk_rows` 和 `::cache_resident` 够到。但 `#[doc(hidden)]` 意味着 crate 不为它们提供任何稳定性承诺。请把它们当内部实现,改去用 [6.2.5](#625-调整阈值公开接口) 里管着它们的那几个旋钮。
下面是各个部分分别在哪里被调用:
- `Dense::forward` 是一次 `gemm_fused` 调用。它把线性乘积、按列的偏置,以及 ReLU 激活,融进了 1 趟里。
- `Dense::backward` 是 2 次普通的 `dot` 调用。第一次算权重梯度。第二次算输入梯度。
- `SimpleRNN`、`LSTM` 和 `GRU` 用 `dot` 把输入一次性投影好。随后每个时间步都用 `gemm_fused` 融合它的递归投影。GRU 往一个更大缓冲区的切片里写时,会降级到普通的 `gemm`。
- im2col 卷积引擎把每个 filter 的偏置融进它的前向 GEMM。它的 2 次反向 GEMM 都走 `dot_par`。当 batch 那一路的扇出已经喂满线程池时,每个样本的乘积就保持串行。
- `LinearRegression`、`LogisticRegression`、`LinearSVC` 和 `SVC` 用 `matvec` 做预测和求梯度。`machine_learning::linalg` 与 LDA 里的幂迭代和单边 Jacobi 迭代也是。LDA 还用 `dot_par` 构造它的散度矩阵。
- PCA、核 PCA、KMeans,以及 `machine_learning::types` 里的核矩阵代码,用的是 `dot`。
- KNN、t-SNE 和 MeanShift 用 `cache_resident` 和 `gemm_chunk_rows`。这两个函数在"逐行 GEMV swarm"与"分块 GEMM"之间为一次成对投影做选择。
只要你调用了上面任何一个模型,就已经在用这套后端了,只不过没直接喊它的名字。
这些乘积是经由一个私有依赖上的 `pub(crate)` 函数走的。你调用不了它们,也不该想着绕过去。RustyML 只重新导出了 gemmkit 的 tuning 模块。你连一个 `Parallelism` 值都无法通过公开 API 叫出名字。
把你的层和估计器搭在公开 API 上,这套后端就白送给你了。要写你自己的线性代数,就改用 ndarray。[6.2.5](#625-调整阈值公开接口) 里那几个旋钮是唯一对外暴露的接口。它们不用重新编译,就能全局改变行为。
## 6.2.2. 为什么不用 ndarray 的 `.dot()`
默认构建下,ndarray 的 `.dot()` 用的是 `matrixmultiply` crate。这是一个纯 Rust 的 GEMM,而且它做得相当不错。你自己写代码时就该用它。但放到一个要调用几百万次、覆盖各种形状的训练循环底下,它就不是正确的选择了。
matrixmultiply 并不是什么朴素的标量内核。它会在运行期依据检测到的 CPU 特性挑选微内核。这些特性在 x86-64 上是 FMA 加 AVX2、AVX 或 SSE2,在 aarch64 上是 NEON。"gemmkit 用了向量化而 `.dot()` 没有"这种说法是错的。真正的差别来自 3 件事:针对形状的专用路径、融合尾声,以及多线程。ndarray 没有打开 `matrixmultiply` 那个可选的多线程 feature,所以这套构建里的 `.dot()` 只跑在 1 个线程上。
**针对形状的专用路径。** gemmkit 不跑单一一套分块算法。它按形状在几条路径之间挑选。这些路径包括一条专门的矩阵-向量路径、一条给浅 `k` 用、干脆跳过打包的原地路径,以及一条给小 `m` 小 `n` 用的路径。单一的通用内核能把这些形状算对,但算得慢。训练循环里满地都是这些形状。
**融合尾声。** `gemm_fused` 会在输出分块还留在寄存器里的时候,就在内核里把按列的偏置和激活函数施加上去。`.dot()` 没有这项功能。同一个 `Dense::forward` 若照着 `.dot()` 写,就要在输出上走 3 趟:乘积、偏置、激活。这条路径只需要 1 趟。[6.2.4](#624-确定性与可复现性) 记录了让这一点变得安全的那条保证:融合的结果与不融合的那串操作逐位相等。
**多线程。** `matrixmultiply` 的 `threading` feature 藏在 ndarray 自己的可选 feature `matrixmultiply-threading` 后面,而 RustyML 没有打开它。gemmkit 会自己开线程,并且是否值得开由它自己判断。[6.2.3](#623-gemmkit-如何调度一次乘积) 完整讲了这个决定。
操作数的步长直接透传给内核。这是一种便利,而不是相对 `.dot()` 的优势,因为 `.dot()` 自己处理步长也毫无问题。适配器接受任何 `S: Data` 的 `ArrayBase<S, Ix2>`。这包括拥有所有权的 `Array2`、`ArrayView2`、转置视图(`a.t()`)、非连续的切片,甚至步长为负的视图。它不会先复制或物理转置任何东西。转置视图不过是交换了一对步长,引擎直接读任意步长。
这一点很关键,因为反向传播里满是转置的操作数。`dot(&input.t(), &grad_upstream)` 就是 `Dense` 里算权重梯度的写法。走 `.dot()` 的路子要么得把它们复制进连续缓冲区,要么会丢掉这种融合的步长处理。`src/math/matmul.rs` 里的测试证实了这一点。`.t()` 操作数和 `s![..;2, ..]` 这样行步长的切片,都会把正确的步长喂给内核。两者都与一个独立的参考乘积吻合。
crate 从前那套手写后端,有 2 件事根本做不到,而 gemmkit 现在做到了。融合尾声是其中第一件。第二件是 gemmkit 的分块与任务顺序不依赖 worker 数量。正是这种独立性,把 [6.2.4](#624-确定性与可复现性) 里的可复现性声明,从一句托辞变成了一个承诺。
## 6.2.3. gemmkit 如何调度一次乘积
RustyML 在每个调用点上只做 1 个调度决定,而且这个决定是二选一的。它传 `Parallelism::Rayon(0)`,意思是"你自己看着办"。或者它传 `Parallelism::Serial`,意思是"这个线程已经在一个 rayon 并行区域里了,别再 fork 一次"。后一种写法值得记住。卷积引擎的反向传播和 MeanShift 的种子循环用的都是它。它关乎的是别 fork 两次,从来无关正确性。
这个选择之外的一切都归 gemmkit 管。这包括串行还是并行、worker 数量、工作跑在哪个池子里,以及这个形状要不要干脆走一条受带宽限制的路线。
一个 gemmkit 旋钮按优先级依次定值。单次调用的实参,比如 `Parallelism` 请求,压过程序里的 `set_*` 调用。`set_*` 调用压过 `GEMMKIT_*` 环境变量。环境变量压过编译期默认值。每个环境变量只在该旋钮首次被访问时读一次,之后整个进程都用缓存值。
`set_*` 调用是无条件写入的。只要进程里有任何东西调过一次 setter,对应的环境变量在这次运行的余下时间里就没有效果了。这正是 RustyML 绝不替你调用 setter 的原因。一个解析不出非负整数的 `GEMMKIT_*` 值,只会在 stderr 上告警一次,然后退回编译期默认值。性能配置文件里打错一个字,绝不会让进程崩溃。
**工作量闸门。** `parallel_threshold` 是串行与并行的切换点。它的默认值是 `48 * 48 * 256`,也就是 589,824。这个闸门比的是 `m * n * k` 的乘积,不是 FLOPs,所以哪儿都没有那个 2 倍系数。请把单位看仔细,因为本页的旧版本比的是 FLOPs。低于这个闸门的问题,无论你请求了多少 worker,都只在 1 个线程上跑。
这一档里正是那些微小的 GEMM:RNN 和 LSTM 的时间步,以及在紧凑循环里被调用的小型 Dense 层。让它们保持串行是正确的选择,而不是偷懒。把工作派发到线程池的开销,本身就盖过了乘法。
**worker 爬坡。** 越过闸门之后,自动路径也不会一下子抓满所有核。`par_mnk_per_worker` 在原生目标上默认是 2,000,000。它规定了每多要一个 worker,还得多带来多少额外的 `m * n * k` 工作量:目标 worker 数是 `mnk / par_mnk_per_worker`,下限为 1,上限受核数和任务数约束。
这条爬坡按工作量而非按维度来,是因为实测的最优点跟随的是总工作量,而不是线性尺寸。gemmkit 自己在 Ryzen 9950X 上做的标定证实了这一点。一个 `128^3` 的乘积(约 2e6)串行跑最快。一个 `192^3` 的乘积(约 7e6)想要 2 或 3 个 worker。一个 `384^3` 的乘积(约 5.7e7)已经想要全部 32 个硬件线程。没有哪一条沿单一维度设的阈值,能同时照顾这条曲线的两端。
**线程池分档。** 本页的旧版本说这套后端不自带线程池。这已经不对了。`pool_classes` 会建起若干持久的、尺寸严丝合缝的私有 rayon 池,分成若干档。这些档从机器宽度的一半开始逐级折半:1 档是 width/2,2 档再加 width/4,3 档再加 width/8。自动算出的 worker 数会贴到仍能容纳它的最小那一档。
理由在于 rayon 的 fork-join 税。这项税跟随的是池子的空闲余量,也就是池宽减去真正在干活的 worker 数,而不是 worker 数本身。8 个 worker 待在一个 8 宽的池子里,会大幅优于同样 8 个 worker 在一个 32 宽的全局池里空转。
这些分档池只建一次,之后热着复用。它们不会每次调用都重建。设成 `0` 就完全禁用它们。默认值按架构分裂:x86-64 上 2 档,aarch64 上 1 档,其余每个目标上 0 档,等着在设备上验证。
如果调用线程本身已经是一个 rayon worker,比如在一个嵌套的 GEMM 里,或者在你自己 install 的池里,gemmkit 会跳过这些档位。它会直接在当前池里跑。这正是这些乘积仍能干净地嵌进外层并行区域的原因。它们不会在你的池子上再摞一个池。
**matvec 自成一个开销类别。** gemmkit 会识别出 `m == 1` 或 `n == 1` 的形状,改走一条专用的、受带宽限制的路径,而不是走通用驱动。`matmul::matvec` 存在的意义,就是把一个 `Array1` 摆成能触发这条路径的 `[k, 1]` 列。这条路径根本不查 `parallel_threshold`。
它在一个字节下限之下保持串行,也就是 `gemv_parallel_bytes`,默认是 `0`(意思是"按缓存大小推导")。推导出来的下限是 1 个核的私有 L2。低于它时,被触及的数据是 L2 常驻的,那个核已经吃满了全部 L2 带宽,再切分只会白白添上 fork-join 开销,也换不回任何 DRAM 带宽。
越过这个下限之后,worker 数会随着被触及的字节数攀一道梯子。梯子的每一级就是通用驱动用的那套精确匹配线程池档位,被触及的字节数每上一个 `gemv_tier_step` 倍,就往上爬一级。所以一个刚刚越过下限的 matvec 拿到的是最窄的那一档,而不是完整的内存并行宽度。`gemv_thread_cap` 会盖掉这道梯子:非零值就是逐字采用的宽度,在任何规模上都钉死不变。两者都默认是 `0` 表示自动。
`gemv_axpy_par_min_rows` 在此之上再加一道跟形状有关的闸门。列主序的 matvec 在输出行数低于这个值时,会把所有行留在 1 个 worker 上,因为那里输出行轴是内存的内层轴,切开它会让每个 worker 都在整个矩阵上做跨步游走。行主序矩阵不受影响,因为它的 worker 各自拥有整条 `k` 连续的行。RustyML 的操作数是行主序的,所以 `matvec` 从不查这道闸门。
最后一个旋钮 `gemv_threshold`,限定向量那一侧最大能到多少,超过就把这个形状退回通用驱动。它的默认值是 `usize::MAX - 1`,实际上等于无上限。所以在实践中,一个 gemv 形状的问题总会走 gemv 路径,除非你自己调低这个旋钮。
这几个之外,还有十来个旋钮:`kc`、`rhs_pack_threshold`、`lhs_pack_*` 一族、`small_k_threshold`、`small_mn_dim`、`prefetch_min_bytes`,以及其他一些。本页不会把它们逐一列出来,因为这样一张表迟早会过期。它们在 gemmkit 自己的 docs.rs 页面上有文档。每一个旋钮都能经由 `rustyml::tuning::matmul::backend` 够到。`gemmkit-tune` 自动调优器会在你的目标机器上替你把它们扫一遍。
有一条注意事项适用于其中每一个旋钮。gemmkit 的参考机是一台 Ryzen 9950X(x86-64)和一台 M4 Max(aarch64)。凡是切换点依赖架构的旋钮,都为每种架构带一个各自独立的默认值,按 `cfg(target_arch)` 分裂。除非另有说明,本页引用的数字都是 x86-64 那一侧的值。
## 6.2.4. 确定性与可复现性
本页的旧版本说结果在同一台机器上可复现,但未必逐位相同。这个说法已经不成立了,请把它丢掉。
那句旧托辞之所以存在,是因为 crate 从前那个按行切分的包装函数,会给每一块不同的 `m`。而内核内部沿 `k` 的分块又依赖 `m`,于是求和顺序会随线程数漂移。那种按行切分已经没了,托辞也跟着没了。`src/math/matmul.rs` 现在写下的是一句实打实的承诺:
> gemmkit 的分块与任务顺序不依赖 worker 数量。所以在固定的机器和固定的配置下,同一个乘积会逐位复现同样的结果,无论是多少个线程跑的。结果也会在多次运行之间重复出现。融合尾声(偏置与激活)与"先做普通乘积、再做同样的标量映射"逐位相同。
这不是一句愿景。模块自己的测试套件把上面每一部分都钉住了。
- `dot_par_thread_count_independent_f64` 拿一个 `96^3` 的形状、一个 `256 x 64 x 64` 的形状,以及一个瘦 `k` 的 `64 x 8192 x 64` 形状。它把每个形状先串行跑一遍,再用 `Rayon(2)`、`Rayon(4)`、`Rayon(8)`、`Rayon(16)` 和 `Rayon(32)` 各跑一遍。它断言 `to_bits()` 在每个分支上都相等。之所以放进那个瘦 `k` 的形状,是因为它最容易诱使实现去做 split-`k` 归约,而那恰恰会打破这条性质。
- `dot_par_thread_count_independent_f32` 对 `f32` 做同样的检查。
- `matvec_serial_and_auto_agree_bitwise` 覆盖那条受带宽限制的 gemv 路径。那里每个输出元素都是在 1 个 worker 上沿整个 `k` 归约出来的。
- `dot_run_to_run_deterministic` 和 `matvec_run_to_run_deterministic` 覆盖同一台机器上的重复调用。
- `gemm_fused_bias_relu_bitwise_matches_unfused` 检查带 `Bias::PerCol` 和 `Activation::Relu` 的 `gemm_fused`,是否与"先做一次普通 `dot`、再做同样的标量加偏置并截断"逐位相等。正是这一点,让把偏置和 ReLU 融进 `Dense` 前向成了一次免费的优化,而不是一笔数值上的交易。
**固定的机器和固定的配置**这几个字仍然承重。不同的 CPU 会挑到不同的 SIMD 宽度,因而是不同的累加布局。改动某个旋钮也可能改变分块。跨机器的逐位相等依然不做承诺,任何多线程 BLAS 也不做这个承诺。
但在同一台机器上的同一个二进制里,worker 数量已经不再是你必须费心推理的变量。对一次你想日后重放的训练运行来说,这才是真正要紧的部分。至于可复现性里播种那一半——权重初始化、打乱、dropout 掩码——见 [7.1. 可复现性与随机种子](../Chapter-07/7.1._可复现性与随机种子.md)。
[6.3. 并行归约](./6.3._并行归约.md) 里的那些确定性归约给出的是更严格的保证。它们从构造上就给出相同的结果,与机器无关,而不只是与 worker 数量无关。
## 6.2.5. 调整阈值:公开接口
这才是你能直接调用的部分,而且它分 2 层。串行与并行的抉择归 [`gemmkit`](https://crates.io/crates/gemmkit) 后端管,[6.2.3](#623-gemmkit-如何调度一次乘积) 已经讲过。crate 从前手写并暴露的那几个按数据类型分的 FLOPs 闸门已经没了。
`rustyml::tuning::matmul` 随 `math` feature 提供,因此也在 `full` 之下。它仍然掌管着调用侧的分块策略。它还重新导出了后端自身的旋钮,好让你永远不必直接依赖 `gemmkit`。
这个重导出走的是 `gemmkit-ndarray`,也就是 RustyML 真正调用的那个适配器,而不是自己再依赖一份 `gemmkit`。如果你无论如何都要把 `gemmkit` 加进自己的 `Cargo.toml`,这一点就很关键:这些旋钮是进程全局的原子量,所以一旦 cargo 把你的 `gemmkit` 解析到跟适配器不同的版本,你拿到的就是第二份副本,在它上面调 `set_*` 对 RustyML 的乘积毫无影响。经由 `rustyml::tuning::matmul::backend` 则不可能落到错的那一份上。
| 函数对 | 默认值 | 控制的内容 |
|---|---|---|
| `get_chunk_elems` / `set_chunk_elems` | 33,554,432 | 分块乘积中 1 个行块的元素预算 |
| `get_cache_resident_max_bytes` / `set_cache_resident_max_bytes` | 67,108,864 | 常驻缓存的尺寸阈值,设为你机器的共享 L3 |
| `matmul::backend::*` | 见 gemmkit | 后端的每一个旋钮,各自对应一个 `GEMMKIT_*` 环境变量 |
`cache_resident_max_bytes` 是你最可能想改动的那个旋钮。把它设为你实际的共享 L3 大小。默认值 64 MiB 是个猜测,它周围那一带没有标定过。
要调串行与并行的切换点,请用 `matmul::backend`。`set_parallel_threshold` 卡的是 `m * n * k` 的乘积。`set_gemv_threshold` 卡的是 matvec 那条路径。后端的每个旋钮也都能从一个 `GEMMKIT_*` 环境变量读取。`gemmkit-tune` 自动调优器能一次性产出整台机器的配置文件,所以你很少需要手工挑数字。
把这些旋钮在启动时、进入热循环之前一次性设好。它们是全局的,作用于整个进程。从 RustyML 这边调用一个后端 `set_*` 函数,会让对应的 `GEMMKIT_*` 环境变量在这次进程的余下时间里失声。这会覆盖掉你本来通过环境变量设好的配置文件。这正是 RustyML 绝不替你去设这些旋钮的原因。
完整的来龙去脉、标定流程,以及这些旋钮如何与归约、逐元素运算的阈值配合,见 [7.3. 性能调优与并行](../Chapter-07/7.3._性能调优与并行.md)。
## 6.2.6. 并行何时划算,以及如何测量
这些阈值把"多线程在哪里有用"编码了进去。`benches/benchmarks/matmul_kernels.rs` 里的形状扫描证实了这一点。用 `cargo bench --bench matmul_kernels` 跑它。`Dense::forward` 是 1 次融合的 GEMM 调用,别无其他。偏置和激活都跑在内核的尾声里,不是额外的趟数。这个扫描测了 6 种形状,标签按 `batch x in_features x out_features` 写,也就是 `m x k x n`:
- 4 级近方形的梯子是 `small_256x256x256`、`medium_512x1024x1024`、`big_1024x2048x2048` 和 `huge_2048x2048x2048`。它们从头到尾走完了整条 worker 爬坡。哪怕最小的那一档,`m*n*k` 也有 `16,777,216`,约是工作量闸门的 28 倍,所以它们没有一个是串行案例。这条梯子展示的是随着工作量增长、爬坡如何多发 worker。它还展示了到了顶端、问题大到想要全宽时,线程池分档如何退出画面。
- `wide_256x256x8192` 是宽 `n` 的情形。这里独立的输出列多得是,工作切分毫无别扭之处。对任何多线程 GEMM 来说,这都是好办的形状。
- `thin_256x8192x256` 是有意思的那个形状。它名字里的"thin",指的是 `k` 的两个邻居瘦,不是整体瘦:`m` 和 `n` 都是 256,而 `k` 是 8192。这是一个深 `k` 的乘积。一种常见的直觉认为凡是细长的形状就一定受带宽限制,但这个形状扎扎实实是受算力限制的:约 1.07 GFLOP 对上约 17 MB 的操作数。它是深度分块决策最要紧的形状,这也是扫描里带上它的原因。
有 2 种情形,这个基准是有意不覆盖的。
真正的 matvec 从不出现在里面,因为一次 `Dense` 前向永远不是 matvec。matvec 会彻底离开通用驱动,转投 gemmkit 的 gemv 路径。它改用一个从 1 个核的私有 L2 推导出来的字节下限来卡,随后随着被触及的字节数攀一道 worker 梯子,因为 DRAM 饱和所需的 worker 数远少于机器的逻辑核数。这条路径受限于带宽,多加的核在那里开始划算的时机,远早于一个受算力限制的 GEMM 能摊平线程派发开销的时候。
闸门以下的乘积同样不在其中:RNN 和 LSTM 的时间步,以及小型 Dense 层。这里最小的形状本来就已经远在闸门之上。要是你剖析一个 RNN,发现 rayon 的开销占了大头,别去调低 `parallel_threshold`。那些乘积本就是按设计保持串行的,开销来自别的地方。
你可以在本机上、围着一个固定的乘积来回拨动闸门,快速比一比串行与并行。下面的例子通过一个公开的 `Dense` 层驱动这套后端,把同一个乘积两种方式各计时一遍。把打印出的数字当草图看,不要当基准,因为单次调用噪声很大。要真实数据,就用上面那个会预热、会重复的 criterion 基准。
```rust
use ndarray::Array;
use rustyml::neural_network::layers::{Activation, Dense};
use rustyml::neural_network::traits::Layer;
use rustyml::tuning::matmul;
use std::time::Instant;
fn main() {
// 一次 Dense 前向就是后端的一次 GEMM:input (batch, in_features) @ weights (in_features, units)。
let (batch, fin, fout) = (256usize, 256usize, 256usize);
let mut layer = Dense::new(fin, fout, Activation::ReLU)
.unwrap()
.with_random_state(42);
let x = Array::from_elem((batch, fin), 0.5f32).into_dyn();
// 后端卡的是 m*n*k 的乘积,不是 FLOPs——没有那个 2 倍系数。
let work = batch * fin * fout;
println!(
"backend parallel gate = {}; this product = {} (parallel: {})",
matmul::backend::parallel_threshold(),
work,
work >= matmul::backend::parallel_threshold()
);
let warm = layer.forward(&x).unwrap();
assert_eq!(warm.shape(), &[batch, fout]);
// 把闸门抬到刚好高于这个工作量,强制这个乘积走串行。
let saved = matmul::backend::parallel_threshold();
matmul::backend::set_parallel_threshold(work + 1);
let t0 = Instant::now();
for _ in 0..20 {
let _ = layer.forward(&x).unwrap();
}
let serial = t0.elapsed() / 20;
// 恢复闸门,让同一个乘积改走并行策略。
matmul::backend::set_parallel_threshold(saved);
let t1 = Instant::now();
for _ in 0..20 {
let _ = layer.forward(&x).unwrap();
}
let parallel = t1.elapsed() / 20;
println!("serial ~ {serial:?} / forward");
println!("parallel ~ {parallel:?} / forward");
}
```
闸门和工作量由默认值和形状固定下来,所以它们会原样打印出来。耗时依机器而定,所以下面这段输出给的是形状和类别,不是固定的数字:
```text
backend parallel gate = 589824; this product = 16777216 (parallel: true)
serial ~ <duration> / forward
parallel ~ <duration> / forward
```
留意第二次调用 `set_parallel_threshold` 恢复保存值这一步。程序里的 setter 会永久盖住对应的 `GEMMKIT_PARALLEL_THRESHOLD` 环境变量。所以像这样的一段代码,即便"恢复"了,也已经把这个旋钮在余下的进程里钉死在代码里。这里无害,因为恢复的正是进程启动时的那个值。但这也是你不该把 setter 撒得满库都是的理由。
在这么小的乘积上,或者在核数不多的机器上,并行没有更快也别意外。这正是闸门存在的全部意义,也是默认值让闸门以下的乘积保持串行的原因。把 `batch`、`fin`、`fout` 放大到基准里那些更大的形状,并行这一支就会反超。
如果你宁愿自己写矩阵乘积,也不想绕经某个层,那就用 ndarray。它被有意留在这套后端之外:
```rust
use ndarray::array;
fn main() {
// RustyML 的 matmul 入口是 crate 内部的;你自己的 matmul 用 ndarray 的 `.dot()`。
let a = array![[1.0_f64, 2.0, 3.0], [4.0, 5.0, 6.0]]; // 2x3
let b = array![[1.0_f64, 0.0], [0.0, 1.0], [1.0, 1.0]]; // 3x2
let c = a.dot(&b); // 2x2
assert_eq!(c, array![[4.0, 5.0], [10.0, 11.0]]);
println!("A.dot(B) shape = {:?}", c.shape());
}
```
把你的模型搭在公开的层和估计器上,gemmkit 就白送给你了。这包括它的运行期 ISA 派发、按工作量调度与线程池分档、融合尾声,以及与 worker 数量无关的数值。什么都不用配置。
当某台特定机器想要不同的切换点时,请用 [6.2.5](#625-调整阈值公开接口) 里的那些阈值。更深入的讲解见 [7.3. 性能调优与并行](../Chapter-07/7.3._性能调优与并行.md)。与这套后端并排的距离内核,见 [6.1. 距离度量](./6.1._距离度量.md)。共享它并行机制的归约,见 [6.3. 并行归约](./6.3._并行归约.md)。