资讯动态

【RustyML入门】6.2. 矩阵乘法

发布时间:2026/8/21 11:00:38 来源:尧图企业网站定制
6.2. 矩阵乘法每一次 Dense 层的前向传播、每一个循环时间步都归结为同一件事一次矩阵乘积。线性模型的预测以及 KNN 和 t-SNE 内部的成对投影也是如此。RustyML 不把它们交给 ndarray 的.dot()。RustyML 把它们委托给gemmkit一个纯 Rust 的 GEMM 引擎。crate 经由它那层零拷贝的gemmkit-ndarray适配器够到 gemmkit。mathfeature 只点名适配器引擎由适配器带进来。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准备数据。不要直接用这套后端。6.2.1. 这套后端是什么为什么它是内部的这里有 2 层分清楚会省很多事。引擎是 gemmkit。它在带步长的视图上计算C - alpha*A*B beta*C。它在运行期挑选自己的指令集。它自己做打包和分块并且掌握全部调度决策。gemmkit-ndarray是一层很薄的适配器。它直接从ArrayBaseS, 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 另外还能在可选的halffeature 下支持f16和bf16在int8feature 下支持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 里管着它们的那几个旋钮。下面是各个部分分别在哪里被调用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 里那几个旋钮是唯一对外暴露的接口。它们不用重新编译就能全局改变行为。6.2.2. 为什么不用 ndarray 的.dot()默认构建下ndarray 的.dot()用的是matrixmultiplycrate。这是一个纯 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 记录了让这一点变得安全的那条保证融合的结果与不融合的那串操作逐位相等。多线程。matrixmultiply的threadingfeature 藏在 ndarray 自己的可选 featurematrixmultiply-threading后面而 RustyML 没有打开它。gemmkit 会自己开线程并且是否值得开由它自己判断。6.2.3 完整讲了这个决定。操作数的步长直接透传给内核。这是一种便利而不是相对.dot()的优势因为.dot()自己处理步长也毫无问题。适配器接受任何S: Data的ArrayBaseS, 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 里的可复现性声明从一句托辞变成了一个承诺。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 个线程上跑。这一档里正是那些微小的 GEMMRNN 和 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/22 档再加 width/43 档再加 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 9950Xx86-64和一台 M4 Maxaarch64。凡是切换点依赖架构的旋钮都为每种架构带一个各自独立的默认值按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. 可复现性与随机种子。6.3. 并行归约 里的那些确定性归约给出的是更严格的保证。它们从构造上就给出相同的结果与机器无关而不只是与 worker 数量无关。6.2.5. 调整阈值公开接口这才是你能直接调用的部分而且它分 2 层。串行与并行的抉择归gemmkit后端管6.2.3 已经讲过。crate 从前手写并暴露的那几个按数据类型分的 FLOPs 闸门已经没了。rustyml::tuning::matmul随mathfeature 提供因此也在full之下。它仍然掌管着调用侧的分块策略。它还重新导出了后端自身的旋钮好让你永远不必直接依赖gemmkit。这个重导出走的是gemmkit-ndarray也就是 RustyML 真正调用的那个适配器而不是自己再依赖一份gemmkit。如果你无论如何都要把gemmkit加进自己的Cargo.toml这一点就很关键这些旋钮是进程全局的原子量所以一旦 cargo 把你的gemmkit解析到跟适配器不同的版本你拿到的就是第二份副本在它上面调set_*对 RustyML 的乘积毫无影响。经由rustyml::tuning::matmul::backend则不可能落到错的那一份上。函数对默认值控制的内容get_chunk_elems/set_chunk_elems33,554,432分块乘积中 1 个行块的元素预算get_cache_resident_max_bytes/set_cache_resident_max_bytes67,108,864常驻缓存的尺寸阈值设为你机器的共享 L3matmul::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. 性能调优与并行。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 n4 级近方形的梯子是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 基准。usendarray::Array;userustyml::neural_network::layers::{Activation,Dense};userustyml::neural_network::traits::Layer;userustyml::tuning::matmul;usestd::time::Instant;fnmain(){// 一次 Dense 前向就是后端的一次 GEMMinput (batch, in_features) weights (in_features, units)。let(batch,fin,fout)(256usize,256usize,256usize);letmutlayerDense::new(fin,fout,Activation::ReLU).unwrap().with_random_state(42);letxArray::from_elem((batch,fin),0.5f32).into_dyn();// 后端卡的是 m*n*k 的乘积不是 FLOPs——没有那个 2 倍系数。letworkbatch*fin*fout;println!(backend parallel gate {}; this product {} (parallel: {}),matmul::backend::parallel_threshold(),work,workmatmul::backend::parallel_threshold());letwarmlayer.forward(x).unwrap();assert_eq!(warm.shape(),[batch,fout]);// 把闸门抬到刚好高于这个工作量强制这个乘积走串行。letsavedmatmul::backend::parallel_threshold();matmul::backend::set_parallel_threshold(work1);lett0Instant::now();for_in0..20{let_layer.forward(x).unwrap();}letserialt0.elapsed()/20;// 恢复闸门让同一个乘积改走并行策略。matmul::backend::set_parallel_threshold(saved);lett1Instant::now();for_in0..20{let_layer.forward(x).unwrap();}letparallelt1.elapsed()/20;println!(serial ~ {serial:?} / forward);println!(parallel ~ {parallel:?} / forward);}闸门和工作量由默认值和形状固定下来所以它们会原样打印出来。耗时依机器而定所以下面这段输出给的是形状和类别不是固定的数字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。它被有意留在这套后端之外usendarray::array;fnmain(){// RustyML 的 matmul 入口是 crate 内部的你自己的 matmul 用 ndarray 的 .dot()。letaarray![[1.0_f64,2.0,3.0],[4.0,5.0,6.0]];// 2x3letbarray![[1.0_f64,0.0],[0.0,1.0],[1.0,1.0]];// 3x2letca.dot(b);// 2x2assert_eq!(c,array![[4.0,5.0],[10.0,11.0]]);println!(A.dot(B) shape {:?},c.shape());}把你的模型搭在公开的层和估计器上gemmkit 就白送给你了。这包括它的运行期 ISA 派发、按工作量调度与线程池分档、融合尾声以及与 worker 数量无关的数值。什么都不用配置。当某台特定机器想要不同的切换点时请用 6.2.5 里的那些阈值。更深入的讲解见 7.3. 性能调优与并行。与这套后端并排的距离内核见 6.1. 距离度量。共享它并行机制的归约见 6.3. 并行归约。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价