复数拆分内核
复数 GEMM 是唯一不搭乘 FloatGemm 的同质类型家族。理论上它本可以搭乘。Complex<f32> 和 Complex<f64> 都在自身类型里累加,所以 Lhs = Rhs = Acc = Out。这正是浮点家族能处理的形态。
问题出在内存布局上。num_complex 把一个复数存成相邻的 (re, im) 对。于是从复数切片加载的 SIMD 寄存器持有的是 re, im, re, im, ...。一次复数乘法需要跨通道组合:实部是 re*re - im*im,虚部是 re*im + im*re。
在交错的通道上,这些组合迫使最内层循环里出现 shuffle 和 fmaddsub 一类指令。这个代价每个深度步都要重复一次,总共 O(mnk) 次。gemmkit 的做法是在打包时把布局改写一次,让热循环保持纯实数 FMA,循环内完全没有 shuffle。
本页依次讲五件事:拆分设计本身、共轭如何从中免费落出、内核经由的接缝、微块形状背后的寄存器预算算术,以及数值上的保证。
拆分布局
家族是 gemmkit/src/kernel/complex.rs 里的 ComplexGemm<T, CONJ_A, CONJ_B>。设计的核心就在它的打包例程里。
pack_planar 把每个微面板按结构数组(SoA)形式铺开。每个深度步,面板先存 width 个实部,紧接着存 width 个虚部。width 对 LHS 是 mr,对 RHS 是 nr,步长互换的方式与共享的 pack_panels 完全一致。
内核于是用普通的连续加载,取到一整个实部寄存器和一整个虚部寄存器。去交错的成本从 kc 内层循环移到了打包这一步。摊销后的成本变成 O(MK + KN),而不是 O(MNK)。
内核只能消费这种平面布局。所以两个操作数永远都要打包。家族设置 FORCE_PACK_LHS = FORCE_PACK_RHS = true,压过驱动层原本基于代价、可能原地读取操作数的决策。
pack_planar 复刻了 pack_panels 的两条写入路径。一条是先导维连续时的直接遍历。另一条是带步长源时的缓存分块转置,这样行主序操作数打包时不会每个元素都错过一次缓存。两条路径写出的面板逐字节相同,只是写入顺序不同。共享框架见打包与工作区。
共轭是打包时的符号翻转
共轭只对虚部取负。打包本来就单独写虚部平面。所以 conj(A)*B 和 A*conj(B) 在热循环里零成本。
设置 CONJ_A 或 CONJ_B 这个 const 泛型,会让打包器在拷贝时对虚部平面取负。这是真正的取负,+0.0 会映射到 -0.0,与 num_complex 的 .conj() 一致。同一个实数 FMA 循环随后原样运行,任何地方都没有逐元素的共轭分支。
这也是强制打包标志存在的第二个理由。当打包做的不止是普通拷贝时,这个变换必须每次都执行。
运行时到编译时的桥梁在 gemmkit/src/dispatch/complex.rs 里。公开入口 gemm_cplx 把 conj_a 和 conj_b 当作普通 bool 接收。run_complex 对这一对布尔值只 match 一次,分发到四个 ComplexGemm 单态化里匹配的那一个。这个分支每次调用只发生一次,绝不会进到循环里面。
这里还带出一处细节。把行主序倾向的 C 规范化的方向交换,实际计算的是 C^T = B^T * A^T。由于 (conj(A)*B)^T = B^T * conj(A)^T,这次交换必须连同共轭标志一起交换。这个交换在 match 之前就完成了。
输出共轭(conjC)没有实现。在退化路径上,也就是 k == 0 或 alpha == 0 时,这些标志根本无关紧要。没有 A*B 项,也就没有什么可共轭的。
热循环:每次复数乘加四条实数 FMA
循环真正跑起来之前,还有一个分层问题要解决。家族是同质的,所以驱动层的约束是 T 取复数类型的 KernelSimd<T, T, T, T>。这个约束只提供 SimdOps<Complex<..>>,而不是拆分内核真正需要的实数运算。
桥梁是 SimdOps::cplx_microkernel 这道接缝。家族的 microkernel 会转发给它。每个 ISA 令牌的覆写,由 gemmkit/src/simd/complex.rs 里的胶水宏 impl_complex_simd! 生成,再转发给唯一一个共享的、ISA 泛型的函数 soa_microkernel,它写在 S: SimdOps<C::Real> 之上。
累加器在家族接缝处保持复数类型。所以复数的 alpha 和 beta 能原样穿过驱动层。但在接缝内部,累加器其实是两组实数寄存器。
薄薄的 SimdOps<Complex<..>> 胶水存在的唯一理由,是让驱动层能读到 LANES,也让同质 blanket 实现能够适用。它的元素运算全是 unreachable!,因为复数 GEMM 从不调用它们。LANES 被设成实数通道数,于是一个实数通道对应一个复数行,驱动层的 mr = MR_REG * LANES 数的正是复数行数。
循环本身,摘自 gemmkit/src/simd/complex.rs:
#![allow(unused)]
fn main() {
for p in 0..kc {
let are_p = a_re.add(p * 2 * mr); // re plane of this depth step
let aim_p = are_p.add(mr); // im plane (offset by `mr`)
let ar: [<S as SimdOps<C::Real>>::Reg; MR_REG] =
core::array::from_fn(|i| simd.loadu(are_p.add(i * lanes)));
let ai: [<S as SimdOps<C::Real>>::Reg; MR_REG] =
core::array::from_fn(|i| simd.loadu(aim_p.add(i * lanes)));
let bre_p = b_re.add(p * 2 * NR);
let bim_p = bre_p.add(NR);
for j in 0..NR {
let br = simd.splat(*bre_p.add(j));
let bi = simd.splat(*bim_p.add(j));
for i in 0..MR_REG {
acc_re[j][i] = simd.mul_add(ar[i], br, acc_re[j][i]); // += ar*br
acc_re[j][i] = simd.fnma(ai[i], bi, acc_re[j][i]); // -= ai*bi
acc_im[j][i] = simd.mul_add(ar[i], bi, acc_im[j][i]); // += ar*bi
acc_im[j][i] = simd.mul_add(ai[i], br, acc_im[j][i]); // += ai*br
}
}
}
}
一次复数乘加是四条融合的实数步骤,流入两组累加器。acc_re 拿到一条 mul_add 和一条 fnma(融合取负乘加,x86 上是 vfnmadd)。acc_im 拿到两条 mul_add。
每条操作都是作用在连续加载和标量广播上的普通逐通道 FMA。在 epilogue 之前,没有任何操作会跨通道。固定的逐 p 顺序是刻意安排的。正是它让同一矩阵的完整微块与边缘微块舍入完全一致。
循环结束后,两组累加器排入平面暂存区。标量 epilogue 折叠复数 alpha(alpha == 1 时跳过复数乘法),按情形合并 beta*C,并在写出时重新交错。这是一次摊销 O(MN) 的扫描,统一处理完整、边缘和带步长的输出微块。
打包中的标量去交错、以及 epilogue 中的标量重交错,都是刻意的选择,不是疏忽。内层循环本身就占了内核总成本的绝大部分。所以在每种 ISA 上,通用的标量路径对这两步来说都是下限。
寄存器压力与 NR 的选择
拆分设计让累加器数量翻倍。一个 MR_REG x NR 的复数微块需要 2*MR_REG*NR 个累加寄存器(一组实部、一组虚部),外加 2*MR_REG 个 A 平面寄存器,以及每个列步 2 个 B 广播。这份预算在 gemmkit/src/dispatch/complex.rs 里逐微块记录在案。它让复数微块比浮点微块更小,本页出现的这些铺块形状也都由它推出。
在 FMA 上,16 个 YMM 寄存器里,c32 取 MR_REG = 1(8 个实数通道对应 8 个复数行),NR = 5。这是 10 个累加器,加 2 个 A 寄存器,加 2 个 B 广播,占 16 个寄存器中的 14 个。空出的那两个很关键。若换成 NR = 6 的 16 占 16 满配微块,会把累加器溢出到栈上,所以 NR 被收缩到 5,代码注释里记下了原因。
AVX-512 的 32 个 ZMM 寄存器缓解了这份压力。c32 取 MR_REG = 2、NR = 6,用掉 24 + 4 + 2 = 32 个中的 30 个。NEON 有 32 个向量寄存器,取 MR_REG = 2、NR = 5,占 32 个中的 26 个,给在途的加载临时量留出空间。wasm 的 simd128 取 MR_REG = 1、NR = 4,共 12 个活跃的 v128 寄存器。
每个 c64 变体都沿用其 c32 兄弟的 MR_REG 和 NR,只是通道数减半。这份预算算术本来就与通道数无关。
精度与可复现性
复数没有任何特殊路径。run_complex 把所有形状都送进 driver::run,没有 gemv、small_mn 或小 k 分支。特殊路径那套机制只服务实数浮点和整数。
这让数值契约很容易陈述。gemmkit/tests/correctness/complex.rs 对每一条都有直接测试。
确定性与线程无关性是按位成立的。分块与线程数无关,所以同一问题的串行与并行运行会产生逐位相同的输出。测试直接断言各线程数下 re/im 的原始位模式相等。
在单次运行内部,四条 FMA 固定的逐步顺序,让同一矩阵的完整微块和边缘微块舍入完全一致。所以结果不依赖于微块边界恰好落在哪里。
共轭完全不引入任何舍入。它只是对精确值的符号翻转。一个专门的小整数输入测试(其中每个乘积与和都精确可表示)用精确相等而不是容差,对照朴素参考检验全部四种共轭组合。
对外部 oracle 而言,标准必然要放宽一些。正确性套件把 gemm_cplx(包括每一种共轭组合)拿去和 gemm crate 比较,用的是套件惯常的 L2 型容差。它还单独检验了一个负行步长视图加共轭的情形,对照一个行反转的参考实现。之所以放宽标准,是因为分块的 SoA 收缩本来就有理由和另一个引擎的求和顺序算出不同的舍入结果。
这正是可复现性契约用在复数上的样子:固定的输入、环境与配置下,结果是同样的位。跨不同引擎或不同求和顺序时,结果只需容差一致。
融合偏置入口是拿到按位保证的例外。gemm_cplx_fused 支持逐行或逐列的复数偏置。它刻意不支持任何激活函数,因为 ReLU 一类基于序的激活在无序域上没有定义。
它的 epilogue 从不触碰内核本身的算术。SoA 内核存入的位,正是普通 gemm_cplx 会存入的那些位。一个局限于单个微块的后处理,只在最后一个深度面板上就地映射这些位。复数家族是 OUT_IS_ACC = true,所以中间面板必须保留原始的部分和。
结果与“先跑 gemm_cplx、再做同样的逐元素偏置加法“逐位相同,对每种形状、每种共轭组合都成立。Identity 单态化会把这个后处理整个常量折叠掉。所以非融合路径不为这个钩子的存在付出任何代价。通用机制见 Epilogue融合。