在自研 AI 加速器上实现 Rejection Sampling 时,框架组已经提供了采用 Gumbel-Max 的 PyTorch 参考实现。算法具备逐 token 并行的结构,但直接翻译成 kernel 后,数据搬运、寄存器依赖和特殊函数延迟仍会让硬件等待。换用快速数学指令,又带来了参考实现中没有的数值问题。
这个项目要解决的问题是:
- 如何把算法的并行空间转化为硬件吞吐
- 修复验证快速路径的采样质量
我先选择适合硬件的数据流,修复验证 低精度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:把主要概率质量移出坏区间
令
精确数学中,这仍等于
一句话地说,由于第一次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 的实现,可以看到分段拟合、中间结果截断和最终舍入都会影响输出。避开显式零平台之后,这些近似误差仍然存在。于是我尝试在数学等价的变换中搜索更好的工作区间:
二次幂缩放主要改变指数;非二次幂缩放还会改变归一化尾数,可能切换拟合分段。因此,不同
初筛采用三个独立 seed、每个 seed 一亿样本;随后对候选方案进行十个 seed、各一亿样本的评估。结果表明,Div10 虽然改善了 KS,χ² 却更差;Div14 在实测候选中取得了更好的综合结果。
| 最终评估指标 | Div8 | Div14 | 变化 |
|---|---|---|---|
| 各 seed 的 KS distance 均值 | −65.03% | ||
| 十亿样本聚合 χ² | 851.50 | 752.78 | −11.59% |
KS 是各 seed 分别计算后取均值;χ² 使用十亿样本的聚合计数。Div14 改善了这些指标,但聚合 χ² 的 p 值仍约为