Keyboard shortcuts

Press or to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

扩展点

gemmkit 的变化点是 trait、const 泛型,以及存放在 OnceLock 槽位里的带类型函数指针,没有宏,也没有 transmute。这套纪律只为了一个目的:库预期会有四类增长,分别是新指令集、新元素类型、新点积指令、新融合变换。每一类都应该以纯增量代码落地,只带来一份简短、可核对的触点清单,并且都不应该改动 driver、打包例程和分块模型本身。

本页把这四份配方展开成实操走查,面向想要扩展这个 crate 本身的人。这些接缝的公开程度足够高,其中最关键的一条,也就是用你自己的内核 family 去驱动泛型 driver,在 crate 之外也同样成立,并且有测试证明这一点。

新的 ISA 后端

一个 ISA 后端是一个零大小的 token,加上一套词汇表实现。wasm simd128 后端(gemmkit/src/simd/wasm.rs)是最近的一个完整范例,值得从头到尾读一遍,因为它只有一个文件加上几行分发代码。

这个 token 唯一的固有行为是 Simd::vectorize,也就是 #[target_feature] 蹦床。运行期 CPU 检测无法与泛型内核上固定的 #[target_feature] 属性搭配使用,所以每次内核调用都跑在一个带注解的小函数里。#[inline(always)] 的原语会折叠进这个函数,让每个 intrinsic 都落在特性已启用的代码生成上下文里:

#![allow(unused)]
fn main() {
// gemmkit/src/simd/wasm.rs
impl Simd for Simd128 {
    #[inline(always)]
    unsafe fn vectorize<R>(self, f: impl FnOnce() -> R) -> R {
        #[target_feature(enable = "simd128")]
        fn inner<R>(f: impl FnOnce() -> R) -> R {
            f()
        }
        inner(f)
    }
}
}

清单如下:

  1. token。 在新的 gemmkit/src/simd/ 模块里加一个 Copy + Send + Sync + 'static 的零大小结构体,按架构做 cfg 门控,并给它加上前面展示的 vectorize 蹦床。

  2. SimdOps<T> 实现。 为该 ISA 加速的每个元素类型都加一份实现,各自需要一个寄存器类型、LANES,以及一套原语词汇:load、store、splat、mul、add、mul_addfnmareduce_sum,如果希望融合浮点 epilogue 能够向量化,还要加上 max/min。这套词汇表刻意做得很“厚“,这样 microkernel 才能始终保持为一个泛型函数。你实现的是原语,而不是内核本身。

    在这里一定要遵守文档化的契约。simd128 的实现用的是 f32x4_pmax,而不是 f32x4_max,因为 trait 里的 max 要求 a 为 NaN 时返回 b,这正是向量与标量 epilogue 之间“ReLU(NaN) = 0“的约定。它还把两个操作数反过来传,写成 f32x4_pmax(b, a),因为 pmax(x, y) 计算的是 x < y ? y : x。如果按自然顺序传参,a 为 NaN 时会返回 NaN,max(-0.0, +0.0) 会返回 -0.0,这两种情况都恰好与契约相反。它还把 mul_add 写成未融合的 mul 后接 add,因为 wasm 没有硬件 FMA,而 relaxed-SIMD 提供的替代方案在规范上是不确定的,会破坏可复现性。

  3. tile 几何形状。 为每个类型选定 (MR_REG, NR),并把它编码成分发模块里各 ISA 包装函数的 const 泛型。这是唯一一个按 (类型, ISA) 设置的旋钮。要明确地做寄存器预算:simd128 对 f32 用 2x4 的布局,也就是 8 个累加器、2 个 LHS 寄存器、1 个 RHS 寄存器,共 11 个活跃的 v128 值,因为 LLVM 的 wasm 后端在活跃向量数超过约 16 个之后就会开始溢出。NEON 则用 4x4 的布局,刻意留出富余的寄存器。

  4. 一个 Dispatched 描述符,以及每条 select_* 阶梯上的一条分支。 记忆化的选择阶梯位于 gemmkit/src/dispatch/ 下:浮点用 select_f32/select_f64,混合精度用 select_f16/select_bf16,整数用 select_i8,复数用 select_c32/select_c64,再加上 map-epilogue 的选择器。每条阶梯分支都把普通、预打包、融合三种入口点和 tile 几何形状捆在一起,所以新增一个 ISA,只需要一个描述符常量,加上每个受益类型一条 match 分支。

  5. 一个 GEMMKIT_REQUIRE_ISA 取值。gemmkit/src/dispatch/isa.rs 里加一个 ForcedIsa 变体及其解析字符串。目前的取值有 scalarfmaavx512favx512vnniavx512bf16neonsimd128auto。遵守“响亮失败“的规则:如果被钉住的 ISA 不受支持,分发必须直接 panic,而不是回退。这样一来,一个原本想测试你这个内核的 CI 任务,就不可能悄悄跑到别的内核上却通过了测试。

  6. 测试基本上是免费搭车的。 tests/simd_conformance.rs 直接构造各个 token,把每个原语拿去和标量参考实现逐一核对。再加上一个 env_isa_* 钉住二进制和一个 CI 任务,分发路线本身也就可测了(见测试与验证)。

你应该完全不需要碰 driver.rs、任何内核 family、pack.rscache.rs。simd128 后端一个都没改。

新的元素类型

元素类型沿着两个小 trait 变化(见标量与内核家族)。Scalargemmkit/src/scalar.rs)只声明恒等常量和累加器类型 Acc。选择 Acc 是这里影响最深远的一个决定,因为它决定了整套舍入方式。f16 选择了 Acc = f32i8 选择了 Acc = i32,这让整数 GEMM 严格精确。KernelFamilygemmkit/src/kernel.rs)则捆起了区分一种运算所需的其余一切:Lhs/Rhs/Acc/Out 类型、打包布局,以及 microkernel。

很多时候根本不需要新的 family。如果新类型只是既有累加器之上的一种窄输入,那就改为在有能力的 token 上实现 KernelSimd<L, R, A, O> 这条拓宽/收窄接缝:拓宽加载,再加一次收窄存储,然后复用泛型 microkernel,就像 MixedGemm<f16>MixedGemm<bf16> 那样。同质情形由一个 blanket 实现覆盖,混合实现不可能与它重叠。真正全新的运算形态,比如平面布局的复数内核,或者重量化的整数 family,才需要拥有自己的 KernelFamily

把新类型接入公开 API,意味着在 gemmkit/src/dispatch/ 下新增一个分发模块,并为该类型准备自己的 OnceLock 槽位。特性检测只跑一次,胜出的单态化入口点会被缓存下来,之后每次调用都只是一次间接调用。如果是两个类型、gemv/small-mn/small-k 的改道逻辑,再加上一点点点积内核选择上的细节,可以照抄 dispatch/mixed.rs 的模式。如果是异质的任务类型,则改为照抄 dispatch/int.rs

这里的开闭性质不是口口相传的说法,而是被 gemmkit/tests/open_closed.rs 直接强制执行的。那个测试定义了 NaiveFloat:一个独立编写、自带打包逻辑和朴素标量 microkernel、只使用公开条目构建的 family。测试用它去驱动完全未改动的公开函数 driver::run,再对照 f64 参考结果做检查。一旦 driver 的改动破坏了 family 这道接缝,这个测试就会连编译都通不过。它同时也是新建一个 family 时可以照着写的模板。

点积指令

vpdpbusdvdpbf16ps 这类指令把好几个深度步骤折进一条操作里,这会重塑累加的舍入方式,所以它们绝不能以“对可移植 tile 循环的巧妙覆写“这种形式出现。为了保持这个区分,这条接缝被拆成了几部分:

  • family 声明 DEPTH_MULTIPLE = Q(大于 1),并经由 pack_kgroup_panelsgemmkit/src/pack.rs)打包。这个函数把 Q 个连续的深度步骤按 lane 交织排列在一起。driver 会把 panel 的深度向上取整到 Q 的整数倍,并保证 k 组不会跨越切片边界。
  • 有能力的 token 覆写 KernelSimd::dot_accumulate,整组整组地消费这些 panel 里的指令组。打包出来的布局是 family 的打包器与覆写它的 token 之间的私有契约。任何符号修正,比如 VNNI 那个带列和补偿的 +128 技巧,都封装在覆写内部完成,这样累加器返回时持有的就是真实的和。
  • SimdOps::accumulate_tile 的覆写只保留给那些调度层面、并且不改变舍入形状的改动,比如需要显式软件流水线的顺序执行核心,或者长度不是编译期常量的可伸缩向量 ISA。它的文档说得很明确:会重塑舍入的指令不在这条接缝的适用范围内,那类指令应该改用带点积接缝的新 family。accumulate_tile 的覆写必须保持确定性,并且要和边缘路径的舍入方式一致。默认实现已经能在任何宽乱序核心上打满 FMA 管线,所以在保留一个覆写之前,先证明它确实值得。

IntGemmVnniBf16DotGemm 是两个现成的例子。IntGemmVnni 对拓宽路径逐位精确,因为整数算术满足结合律。Bf16DotGemm 则改为按容差把关,仍然落在可复现性契约之内。点积内核与深 K 孪生对两者都有深入讨论。

新的融合变换

一个融合变换就是一份 Epilogue 实现(gemmkit/src/kernel/epilogue.rs)。driver 的 last_k 管道、零开销的 Identity 默认值,以及穿过每条特殊路径的路由,全都是免费获得的(见 Epilogue 融合)。真正需要设计的是选定一条应用路径,并遵守一条硬规则:向量路径与标量路径必须逐位一致。完整 tile 走向量路径,边缘和跨步 tile 走标量路径,而同一个输出矩阵可以自由混用这两条路径。

  • 作用在 Acc 类型值上、有自然寄存器形态的变换,设 VECTOR = true 并实现 apply_reg。这是 FusedEpi 的模式。这里要留意 NaN 和带符号零的语义,比如 LeakyRelu 在两种形态下都写成完全相同的 max + slope*min 组合。
  • Acc 收窄成不同 Out 的变换,设 VECTOR_STORE = true 并实现 apply_store。这是 KRequantize 的模式,需要逐情形论证逐位相等,并用一致性扫描把它钉住。
  • 没有划算向量形态的变换,让两个标志都保持 false,一切都经暂存走标量 apply,这对任何 tile 形状都正确。但如果标量值可能与快路径的融合存储相差 1 个 ULP,就改为借用 MapEpi 的技巧:设 VECTOR = true,把 apply_reg 实现成“排空到栈、再逐 lane 应用“,这样这个变换看到的永远是普通 gemm 本会存储的那些精确位。

不管走哪条路径,都要在 gemmkit/tests/epilogue/ 里,紧挨着已有的测试,为这个新变换补上它自己的“gemm 再映射“等价测试。那套测试正是逐位契约真正被强制执行的地方。