国产芯片算子优化:Rejection Sampling
Laceprndpm

在自研 AI 加速器上实现 Rejection Sampling 时,框架组已经提供了采用 Gumbel-Max 的 PyTorch 参考实现。算法具备逐 token 并行的结构,但直接翻译成 kernel 后,数据搬运、寄存器依赖和特殊函数延迟仍会让硬件等待。换用快速数学指令,又带来了参考实现中没有的数值问题。

这个项目要解决的问题是:

  1. 如何把算法的并行空间转化为硬件吞吐
  2. 修复验证快速路径的采样质量
    我先选择适合硬件的数据流,修复验证 低精度log 异常,再通过软件流水隐藏依赖。固定全接受测试点的 kernel cycles 从 1,530,302 降至 128,453,约 11.9×。下面围绕这几个决策展开。

实现还区分 Exact 与 Fast 两种路径:前者需要对齐框架的随机数序列要求;后者不要求同 seed 逐位复现,可以使用硬件 RNG,但仍需验证输出分布的质量。本文的性能与数值优化主要围绕 Fast 路径展开。

1. 算法与语义:需要优化什么?

Speculative Decoding 先由 draft 模型提出候选 token,再由 target 模型验证。记 target、draft 的概率分布分别为 ,当候选被拒绝后,需要从以下权重对应的分布重新采样:

对于存在正权重的情况,Gumbel-Max 可以把这次采样写成:

零权重对应的 score 为负无穷;归一化只会给所有 score 增加同一个常数,因此可以省略。原理与证明可参考 A Review of the Gumbel-max Trick,这里关注它对实现的意义:每个 token 独立打分,最后做一次 argmax,无需构造 CDF。

2. 瓶颈与数据流:为什么读两遍 BF16?

对于sum-exp,当时提出了3种方案:

方案 做法 主要代价
FP32 中间结果落地 第一遍 max unpack 后保存 FP32,第二遍直接读取 增加 FP32 片上写入与读取
两遍 packed BF16 两遍分别读取、unpack,不保存整行 FP32 重复读取 packed 数据与 unpack
Online sum-exp 一遍维护动态最大值和指数和 增加 rescale 计算及循环依赖
这里的读取发生在片上 VMEM 与寄存器之间。输入已经搬到片上,两遍扫描并不意味着从 HBM 重复搬运两遍。

FP32 落地能省掉第二次 unpack,却会增加片上 Load/Store 的负担。Online sum-exp 能省掉第二遍读取,但每次最大值变化都需要调整已有的指数和,增加计算和循环携带的依赖。究竟哪一种更好,需要看硬件上哪类资源先耗尽。最后通过结合硬件后端资源,我选择了第二种。

后续证明高估了FP32路径的LoadStore资源消耗量,原因是packed数据的layout为 [0 128] [1 129]… [127 255] [256 384]所以如果用通用头文件做BF16转FP32为了避免存储时bank-conflict会需要额外两次存储操作,但在sum-exp里理论上可以通过改layout,此时FP32中间落地会有越1cycle/packed bf16的优势。具体Layout可以查看附录。

3. 软件流水:如何把独立工作交叠起来?

数据流确定之后,还需要解决依赖等待。以 sum-exp 为例,每组数据要经过加载、unpack、减最大值、乘换底常数、发起指数计算、取回结果和累加。若一组做完再开始下一组,即使每条指令本身很快,依赖之间也会留下大量空闲时间。

这里要区分两种等待:普通算术指令的结果写回带来较短的 RAW 依赖;异步特殊函数单元则需要更长时间才能产生可取回的结果。发射吞吐和结果可用时间是两个不同约束。 单元能接受下一次请求,不代表上一次结果已经能被消费。

我用四组独立寄存器状态拉开短依赖的复用距离,同时让不同迭代处在不同计算阶段:

逻辑阶段 当前推进的工作
Unpack 已 unpack 的数据减去行最大值
Multipy 乘以换底常数,为指数计算做准备
Pow 向特殊函数单元发起指数计算
Pop 弹出上一轮单元的指数运算,并且完成累加

之所以这么切,一方面ResMII只有10cycles,单个body内无法完全隐藏Pow的13cycle延迟。而对于Unpack,因为Pow操作有1 cycle的hold-issue,如果不拆成2条独立链,没有足够的独立指令来避免连续Pow产生的Pipeline-stall。

每个阶段有四组 worker。一轮同时推进不同迭代的阶段,之后轮换寄存器的逻辑角色;更早发起的异步结果通过 FIFO 取回并累加。

角色轮换可以通过硬件旋转,或通过代码展开实现即,下一轮改变同一批物理寄存器的用途。四组 worker 则是在当前依赖距离和寄存器预算下的选择(具体地对于该硬件的FP操作延迟,即从发射到写回,小于4 cycles);继续增加并行状态会延长值的生命周期,可能引入寄存器压力,不能无限加深。

最终,sum-exp 形成了 10-cycle 的静态循环调度,同时交叠指数请求、算术运算、结果累加、unpack 和加载补充,实测 trace 中该稳态段没有 stall。Residual-Gumbel 的依赖更复杂,28-cycle 的调度体仍约有 1 cycle stall/pass。分别检查这些局部行为,才能知道总周期下降来自哪里,以及哪里仍有空间。

调度前还可以先缩短计算图。例如 argmax 的 value 更新可以由 greater → select 改成独立的 max,比较结果只用于更新 index。这样能缩短 value 的依赖链,但必须同时检查相同最大值和 NaN 的处理语义。

4. 数值修复:从异常定位到 Div8

输出没有 NaN 或 Inf,并不意味着采样质量正常。 Fast 路径使用快速硬件 log 生成 Gumbel 噪声,这里遇到的问题是:log 在输入接近 1 的一个区间内,直接返回 0。

实现采用 base-2 形式,概率项也统一换底,保持 argmax 的数学等价关系。噪声的计算形式为:

先保护第一层,再定位第二层

第一层 log 的输入可能接近 1。我用 缩放输入,再补偿常数:

缩放后的输入上界低于已观察到的坏区间。这与直接 clamp 不同:clamp 会把一段概率质量压到同一个点,而缩放在精确数学中可以通过常数补偿恢复。

但第一层修复之后,1,048,576 个样本仍出现 8627 个精确零。为了定位原因,我把 Gumbel 计算从完整 kernel 中隔离出来,对同一随机输入同时观察中间值与高精度 reference。

当时有两个假设:第一次 log 的普通量化把结果集中到某个值,或第二次 log 自己命中了零平台。根据第一次 log 的量化区间估算,百万样本预计只会命中约 0.5 次,无法解释实际的 8627 次。

逐值 trace 随后显示:某个样本经过第一层 log 和补偿后,第二层输入约为 1.0000002384,恰好落入零平台。按照该平台区间估计的命中数约为 8609 次,与实测量级吻合。概率估算和逐值追踪共同指向了第二层 log。

Div8:把主要概率质量移出坏区间

,第二层采用:

精确数学中,这仍等于 ;硬件实际处理的输入域却发生了变化。原本 附近的高频命中,被移动到 附近,而该区域的概率质量明显更小。除以 8 还可以通过二进制缩放完成,补偿常数也很简单。

一句话地说,由于第一次log后的分布是无界分布,所以不能通过缩放来避开0平台,但是可以通过调整,不伤害精度的情况下,使得0平台质量分布较小。

验证项 Div8 结果
样本数 1,048,576
精确零 8627 → 0
KS
64-bin χ²
同输入绝对误差 P99.9

在这个规模与检验设置下,修复后的结果通过了统计检验。模拟器与实卡的同位置对照中,只有 57 个 FP32 输出不逐位一致,最大绝对差约为 ,支持了基于模拟器进行诊断的有效性。

5. 结果

以下是 BF16、词表大小 154,880、Sequential 模式下的完整 kernel cycles。S、R 沿用测试用例的参数记号;两条路径采用各自基线,不能把加速比直接互相比作优化优劣。

场景 测试参数 基线 cycles 优化后 cycles 加速比
全接受 S=4,R=8 1,530,302 128,453 11.91×
首步拒绝 S=4,R=2 55,896 38,032 1.47×

全接受对照的是初始语义正确实现,首步拒绝对照的是已经包含部分优化的专项基线。这些数字描述固定测试点的 kernel 收益,不代表整套推理服务的加速比。

性能侧用局部调度、trace 和总周期互相验证;数值侧分别检查同输入误差、输出分布和模拟器与实卡的一致性。Div8 的结论是指定规模下通过验证,后续我又主动扩大了验证范围。

Bonus:推进 Div8 → Div14

完成 Div8 修复后,我想继续确认更大样本下的随机质量。将规模扩大到十亿样本,64-bin χ² 达到约 851.50,检出了此前百万样本检验未能识别的系统偏差。

进一步检查快速 log 的实现,可以看到分段拟合、中间结果截断和最终舍入都会影响输出。避开显式零平台之后,这些近似误差仍然存在。于是我尝试在数学等价的变换中搜索更好的工作区间:

二次幂缩放主要改变指数;非二次幂缩放还会改变归一化尾数,可能切换拟合分段。因此,不同 在精确数学中等价,在快速 log 上却可能产生不同的误差分布。

初筛采用三个独立 seed、每个 seed 一亿样本;随后对候选方案进行十个 seed、各一亿样本的评估。结果表明,Div10 虽然改善了 KS,χ² 却更差;Div14 在实测候选中取得了更好的综合结果

最终评估指标 Div8 Div14 变化
各 seed 的 KS distance 均值 −65.03%
十亿样本聚合 χ² 851.50 752.78 −11.59%

KS 是各 seed 分别计算后取均值;χ² 使用十亿样本的聚合计数。Div14 改善了这些指标,但聚合 χ² 的 p 值仍约为 ,极端上尾检查也仍显示不足。它提高了快速路径的随机质量,没有彻底消除偏差;Gumbel-Max 最终取最大值,尾部表现还值得结合实际 token 采样继续验证。