打包与工作区
微内核对输入的要求只有一种形状:对每个深度步,它都要 mr 个 A 元素在内存中连续排列,nr 个 B 元素在内存中连续排列,一条面板接一条面板,中间不留任何空隙。用户矩阵几乎从来都不是这个样子。它们有着任意的行、列步长,尾部除不尽微块尺寸,深度方向的遍历甚至可能跨越内存页。
打包(packing)就是弥合这道缝隙的那一次拷贝。它把每个宏块一次性重排成微面板优先(micropanel-major)布局,让最内层循环每次都能从 64 字节对齐的暂存区里读到纯单位步长的数据流,读到完整的 mr/nr 向量。这次拷贝的开销是 O(mc*kc),而复用它的计算开销是 O(mc*kc*nc)。这就是它能摊销掉的原因,也是驱动器在复用程度低到不值得付出这个开销时会直接跳过它的原因。
一个例程,两个操作数
这次机械的拷贝集中在一个例程里:pack_panels(gemmkit/src/pack.rs,L1 层)。LHS 和 RHS 的布局其实是同一种布局,只是从两个不同的方向去看。LHS 宏块打包成若干条 mr 行高的面板,按列逐列存放:面板 0 存放第 0..mr 行,每个深度步的 mr 个元素连续排列;面板 1 存放第 mr..2*mr 行,以此类推。RHS 宏块打包成若干条 nr 列宽的面板,按行逐行存放。
两种布局都是“每个深度步 width 个连续的前导元素“。唯一的区别在于哪条矩阵轴充当前导轴。于是两个 KernelFamily 钩子调用的是同一个例程,只是把步长对调了一下:
#![allow(unused)]
fn main() {
// gemmkit/src/kernel/float.rs
#[inline]
unsafe fn pack_rhs(
dst: *mut T,
src: *const T,
rs: isize,
cs: isize,
kc: usize,
nc: usize,
nr: usize,
) {
// RHS panels are `nr` columns wide, stored row-by-row: the "leading"
// direction is columns (stride `cs`) and the "depth" is rows (stride
// `rs`), the transpose of the LHS case, handled by swapping strides
unsafe {
pack_panels(
dst, src, /*lead*/ cs, /*depth*/ rs, /*n_lead*/ nc, kc, nr,
)
}
}
}
pack_lhs 是它的镜像:lead = rs、depth = cs、width = mr。当块的大小除不尽时,这个例程会把尾部面板里空出来的车道填零。这样内核就总能读到完整的 mr/nr 向量,乘法本身也就不需要对边缘块做任何掩码处理。
例程内部有两条路径,写出的结果逐字节一致。第一种情况是前导维连续(lead == 1),对应列主序的 A 或行主序的 B。这时每个深度步的 live 个元素在源数据里本来就相邻。这条面板于是就是一串直白的 copy_nonoverlapping 调用,加上尾部补零。
第二种情况是前导维带步长。这时朴素的逐元素收集每读一个元素都可能撞上一次缓存未命中(每个深度步要做 width 次跨步加载)。例程转而跑一趟缓存分块转置:它沿源数据的连续维,以 GEMMKIT_PACK_TRANSPOSE_TILE 个深度步(默认 16)为一条,把每一条散布进面板。这样产出的打包字节和一次纯粹的重排拷贝完全一样,但对跨步的数据源要便宜得多。这正是行主序 A 布局的代价并不比列主序高多少的原因。
点积家族(i8 的 VNNI、bf16 的 vdpbf16ps)还有一个姊妹例程 pack_kgroup_panels。它在此基础上,把每条车道连续的 DEPTH_MULTIPLE 个深度步交织在一起,让一条点积指令能一次吞下整组数据。这种布局属于点积内核与深K孪生的内容。
打不打包,由驱动器决定,而这两个操作数并不对称。
微内核以 mr 宽的向量来读取 A,所以只要 A 的行不是单位步长,或者行面板不完整,驱动器就必须打包它。除此之外,当每个 worker 的列复用超过 GEMMKIT_LHS_PACK_THRESHOLD(aarch64 上默认 256 列,其他平台默认 1024 列)时,驱动器也会打包 A。
对于列主序的 A,驱动器还会在它的深度遍历同时满足步长达到页级、跨度足够宽、并且被足够多列块复用、值得付出这个成本时打包它。以下三个条件必须同时成立:
- 每一步的步长达到半个内存页(
GEMMKIT_LHS_PACK_STRIDE,从Machine记忆化的页大小自动推导)。 - 整条深度切片的遍历(
csa * sizeof(Lhs) * kc)达到GEMMKIT_LHS_PACK_SPAN字节(自动值:4 MiB)。 - 至少有
GEMMKIT_LHS_PACK_REUSE条nr宽的列块复用每一条打包好的面板(min(n, nc) / nr,向上取整,x86 上默认 128,aarch64 上默认 4)。
每一道门槛排除的都是打包不划算的一种情形。一个页级的步长如果发生在仍然驻留缓存的跨度之内,那就只是在重新遍历本来就还热着的缓存行,就地读取 A 的代价反而比付出一次打包更低。跨度这道门槛让 A 保持就地,直到这趟遍历真的宽到足以打垮 TLB,无论后面复用多少次都是如此。复用门槛针对的是另一种失衡:一个瘦高的形状(m 远大于 n)只靠极少的列块就会堆出很大的跨度,把一次昂贵的拷贝摊到太少的复用上并不划算。
复用门槛在不同架构上取值不同,是因为打包和就地读取之间的权衡本身就因架构而异。在 x86 上,打包相对就地读取的代价更高,所以驱动器要等到复用足够多才愿意付出这个代价,默认门槛是 128 条列块。在 aarch64 上,打包相对就地跨步读取的代价更低,驱动器可以更早就选择打包,默认门槛是 4 条列块。
B 则不同,它永远只以广播单个元素的方式被读取,所以任何布局不打包也能用。驱动器打包 B 纯粹是为了复用:每个深度切片打包一次,条件是 m 超过 GEMMKIT_RHS_PACK_THRESHOLD(默认 2048),并且会有足够多的行块反复读取它。由谁来执行这些打包、打包和计算之间的屏障如何安排,属于调度层面的问题,并行执行详细讲述了这部分内容。
预打包操作数
当同一个操作数在一次次调用中反复出现时,每次调用都重新打包就是白费功夫,这正是推理场景的模式:固定的权重反复对上一串流动的激活值。gemmkit/src/api/packed.rs 里的预打包入口会把整个操作数一次性提前打包好。prepack_rhs 沿着任意布局的 B 的步长遍历它,返回一个 PackedRhs<T>。gemm_packed_b 随后就用它来做乘法,完全跳过按调用的 RHS 打包。这套 API 的使用方式见预打包操作数。从架构角度看,有三条性质值得关注。
第一,缓冲区记录了它构建时所用的分块几何:nr、kc、nc。消费它的调用会原封不动地读回这份几何信息。驱动器用记录下来的 kc 和 nc 顶替自己模型算出的结果,只有 mc 仍然按真实的 m 推导。因此,即便打包和消费之间某个调优旋钮发生了变化,面板地址也始终与打包时保持一致。几何本身是通过和普通调用相同的 blocking 模型求解出来的,只是用了一个 tiny_block_dim() + 1 的哨兵行数,让它永远走不到小矩阵分支,因而与最终的 m 无关。
第二,布局只有一个事实来源。prepack_rhs 通过 driver::pack_rhs_full 来填充缓冲区,这个函数铺设面板的顺序,和驱动器自己按片打包时写出的顺序完全一致:最外层是 jc 块,然后是深度切片,再然后是每个切片里 nr 宽的面板。预打包出的字节因此和按调用打包出的字节完全相等,所以在相同配置下,预打包 GEMM 会复现一次普通的 gemm 调用。文档中标注的例外是极小的乘积(m 和 n 都不超过 tiny_block_dim)和 gemv 形状的乘积。普通的 gemm 会把它们改道到特殊路径上,所以它们可能在最后一个 ULP 上有所出入。
第三,缓冲区在整个 GEMM 期间都是只读的,所以每个 worker 都能无需任何同步地共享它。这一点和按调用打包的 B 不同,它不需要任何屏障。
PackedLhs 几乎不用额外的代码,靠的是这套引擎在 A、B 之间的对称性。一个 m x k 的 LHS,本身就是转置乘积 C^T = B^T*A^T 的 RHS。所以 prepack_lhs 只是把步长对调之后,委托给 prepack_rhs_unchecked,gemm_packed_a 也就通过这个转置后的问题来消费它。这种对称性也解释了取向方面的断言:预打包的 B 要求 C 近似列主序(|csc| >= |rsc|),预打包的 A 要求 C 近似行主序。换一种取向的话,分发层就会交换两个操作数的角色,烘焙好的布局就派不上用场了。
int8 feature 增加了一个异构的孪生入口,prepack_rhs_i8 和 gemm_i8_packed_b,有三处刻意做出的不同。
第一,它的布局固定为本进程记忆化的分发所选中的那个整数内核所用的布局:要么是 VNNI 的 k 四元交织布局,要么是加宽内核的普通面板布局。消费入口永远运行同一个家族,所以缓冲区永远不可能被读错。
第二,它把缓冲区深度向上取整到点积内核的 DEPTH_MULTIPLE = 4,并把整个收缩打包成单独一个深度切片,满足驱动器对深度补齐家族的单切片约束。
第三,它刻意绕开了普通 gemm_i8 在低于 GEMMKIT_I8_VNNI_MIN_PAR_MNK 时才会启用的小规模并行加宽回退。vpdpbusd 的缓冲区是四元交织的,加宽内核根本无法消费它。由于整数累加是精确的,无论走哪条路径,结果都与普通的 gemm_i8 逐位一致。
预打包正是在这条路径上收益最大:VNNI 的 RHS 打包本来每次调用都是强制性的,所以在 m 较小时,按调用付出的 O(k*n) 打包开销会压过 O(m*k*n) 的计算开销。
工作区
所有这些打包都需要暂存内存,Workspace(gemmkit/src/workspace.rs)就是这份内存的分配器:一块可以增长、64 字节对齐(足以满足 AVX-512 存储的要求)的缓冲区,按 2 的幂增长,并且从不收缩。每次调用中,Workspace::regions 都会把它切分成 a_regions 份大小相等的 LHS 区域,外加一份共享的 RHS 区域,每份区域都向上取整到对齐边界。
LHS 区域的数目,在按 worker 打包的路径上等于 worker 数,在共享 A 的路径上等于行块数,两种路径下切分方式完全一样。当两个操作数都不需要打包时,驱动器会干脆跳过这次预留,让一个完全就地计算的负载永远不会撑大这个池子。
在字节乘积处失败即拒绝
尺寸计算这部分藏着一个内存安全方面的微妙之处。gemmkit 接受广播(零步长)视图:它们只需要一小片后备存储就能通过边界校验,却呈现出逼近 isize::MAX 的逻辑维度。于是,用来计算打包缓冲区大小的乘积就真的有可能让 usize 溢出,而一旦尺寸回绕成一个偏小的值,就会导致缓冲区分配不足,随后打包时就会越界写入。
驱动器用 checked_mul 守住了元素计数的乘积,但只检查元素计数是不够的。以混合精度路径上的 k = 2^56 为例,那里 kc == k。一个 mc * kc 元素的 LHS 区域,比如说 32 * 2^56 = 2^61,完全能装进 usize,能顺利通过每一项元素级别的检查。可一旦拿这个元素数去乘元素大小、再向上取整到 64 字节对齐,数值就回绕了。溢出恰恰只在元素到字节的换算这一步才会显现,所以守卫也必须设在这里,设在每一份区域大小都要流经的这个咽喉位置:
#![allow(unused)]
fn main() {
// gemmkit/src/workspace.rs
fn region_bytes(elems: usize, esize: usize) -> usize {
elems
.checked_mul(esize)
.and_then(|b| b.checked_next_multiple_of(ALIGN))
.unwrap_or_else(|| workspace_too_large())
}
}
Workspace 会检查每一步:字节乘积、对齐取整、区域总和,以及最后 A + B 的总量。任何一步溢出,都会以和驱动器自身尺寸检查相同的“too large”契约触发 panic。这就是失败即拒绝:代码会大声拒绝一个荒谬的问题,而不是悄悄败坏内存。驱动器无条件地运行这些元素计数守卫,也是出于同样的理由,即便某条路线最终什么都不打包也不例外:一旦跳过这些检查,就等于同时跳过了中止,会把那个荒谬的 k 送进就地循环里,让它近乎无限期地空转下去。
池子、_with 与 no_std
调用者很少会直接看到一个 Workspace,因为一个线程本地的池子会透明地提供一个。常规的 gemm 调用每个线程最多分配一次,之后每次调用都复用同一块缓冲区。
这个池子同时是可重入安全的。嵌套的 rayon 有可能在一个已经身处某次 GEMM 中的线程上再进入一次 GEMM:比如一个 worker 在自己的 for_each 里阻塞时,窃取了另一次 GEMM 的任务;又或者一个批量并行的 worker 内联地跑了其中一个元素。这种情况下,池子的 RefCell 已经被借出,于是 with_thread_pool 会为这一次调用单独发放一块全新的暂存工作区,而不是让程序 panic。打包缓冲区在调用之间不携带任何结果状态,所以这个回退是完全无感的,唯一被跳过的只是那一次的缓冲区复用。
如果需要显式控制,还有 *_with 这一层:每个入口都有一个变体(gemm_with、gemm_packed_b_with 等等),可以传入一个调用者自己持有的 Workspace。从第二次足够大的调用开始,这样做能做到零堆分配,是热循环中大量小乘积、以及延迟敏感代码的合适工具,而 Workspace::with_capacity 甚至能免去首次调用时的分配尖峰。
没有 std 时不存在线程本地存储,所以 with_thread_pool 只会为每次调用简单地新建一个工作区。因为 parallel 本身就依赖 std,这种构建下也就没有线程需要重入。想要复用的调用者可以自己持有一个 Workspace,用 *_with 系列接口,这也是 no_std与WebAssembly 推荐的用法。