扩展点
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)
}
}
}
清单如下:
-
token。 在新的
gemmkit/src/simd/模块里加一个Copy + Send + Sync + 'static的零大小结构体,按架构做cfg门控,并给它加上前面展示的vectorize蹦床。 -
SimdOps<T>实现。 为该 ISA 加速的每个元素类型都加一份实现,各自需要一个寄存器类型、LANES,以及一套原语词汇:load、store、splat、mul、add、mul_add、fnma、reduce_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 提供的替代方案在规范上是不确定的,会破坏可复现性。 -
tile 几何形状。 为每个类型选定
(MR_REG, NR),并把它编码成分发模块里各 ISA 包装函数的 const 泛型。这是唯一一个按(类型, ISA)设置的旋钮。要明确地做寄存器预算:simd128 对f32用 2x4 的布局,也就是 8 个累加器、2 个 LHS 寄存器、1 个 RHS 寄存器,共 11 个活跃的v128值,因为 LLVM 的 wasm 后端在活跃向量数超过约 16 个之后就会开始溢出。NEON 则用 4x4 的布局,刻意留出富余的寄存器。 -
一个
Dispatched描述符,以及每条select_*阶梯上的一条分支。 记忆化的选择阶梯位于gemmkit/src/dispatch/下:浮点用select_f32/select_f64,混合精度用select_f16/select_bf16,整数用select_i8,复数用select_c32/select_c64,再加上 map-epilogue 的选择器。每条阶梯分支都把普通、预打包、融合三种入口点和 tile 几何形状捆在一起,所以新增一个 ISA,只需要一个描述符常量,加上每个受益类型一条 match 分支。 -
一个
GEMMKIT_REQUIRE_ISA取值。 在gemmkit/src/dispatch/isa.rs里加一个ForcedIsa变体及其解析字符串。目前的取值有scalar、fma、avx512f、avx512vnni、avx512bf16、neon、simd128和auto。遵守“响亮失败“的规则:如果被钉住的 ISA 不受支持,分发必须直接 panic,而不是回退。这样一来,一个原本想测试你这个内核的 CI 任务,就不可能悄悄跑到别的内核上却通过了测试。 -
测试基本上是免费搭车的。
tests/simd_conformance.rs直接构造各个 token,把每个原语拿去和标量参考实现逐一核对。再加上一个env_isa_*钉住二进制和一个 CI 任务,分发路线本身也就可测了(见测试与验证)。
你应该完全不需要碰 driver.rs、任何内核 family、pack.rs 或 cache.rs。simd128 后端一个都没改。
新的元素类型
元素类型沿着两个小 trait 变化(见标量与内核家族)。Scalar(gemmkit/src/scalar.rs)只声明恒等常量和累加器类型 Acc。选择 Acc 是这里影响最深远的一个决定,因为它决定了整套舍入方式。f16 选择了 Acc = f32。i8 选择了 Acc = i32,这让整数 GEMM 严格精确。KernelFamily(gemmkit/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 时可以照着写的模板。
点积指令
vpdpbusd、vdpbf16ps 这类指令把好几个深度步骤折进一条操作里,这会重塑累加的舍入方式,所以它们绝不能以“对可移植 tile 循环的巧妙覆写“这种形式出现。为了保持这个区分,这条接缝被拆成了几部分:
- family 声明
DEPTH_MULTIPLE = Q(大于 1),并经由pack_kgroup_panels(gemmkit/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 管线,所以在保留一个覆写之前,先证明它确实值得。
IntGemmVnni 和 Bf16DotGemm 是两个现成的例子。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 再映射“等价测试。那套测试正是逐位契约真正被强制执行的地方。