LLM 推理确定性工程深度实战:从浮点非结合性、Batch 不变性到可复现 Kernel 的全链路解法

执行摘要:过去三年业界在 LLM 推理上的所有工程努力,几乎都指向两个方向——更低的延迟和更高的吞吐。PagedAttention、Continuous Batching、投机解码、量化、KV Cache 分层,无一例外。但有一条暗线长期被忽略:同一个 prompt、同一个模型、同一份权重,在同一个服务上跑两次,输出可能不一样;甚至一次请求的输出会因为队列里挤进了多少其他请求而改变。这不是 bug,而是浮点非结合性与 GPU 归约顺序耦合的必然结果。本文拆解确定性的三个破坏点(浮点归约顺序、kernel 的 split-k/atomic 累积、动态 batch 组成),给出可复现的验证脚本、batch-invariant kernel 的实现要点、CUDA 层确定性归约的写法,以及一份生产落地检查清单。核心结论:推理确定性不是"把 seed 固定住"这么简单,它要求整个 kernel 层的归约顺序与 batch 形状解耦;代价是 10%—30% 的性能,收益是 RL 训练-推理一致、评测可复现、线上问题可回归定位。

一、先复现问题:一句话证明推理不确定

不用复杂的服务框架,纯 PyTorch 就能暴露问题。下面这段脚本对同一批输入分别用 batch=1 和 batch=8 跑同一个线性层,比较两条路径的结果:

import torch

torch.manual_seed(0)
dev = "cuda"
W = torch.randn(4096, 4096, device=dev, dtype=torch.bfloat16)
x = torch.randn(8, 4096, device=dev, dtype=torch.bfloat16)

# 路径 A:逐条推理(batch=1,模拟单请求低负载)
outs_a = torch.stack([torch.matmul(x[i:i+1], W) for i in range(x.shape[0])])

# 路径 B:整批推理(batch=8,模拟高负载合批)
outs_b = torch.matmul(x, W)

diff = (outs_a.float() - outs_b.float()).abs()
mismatch = (outs_a != outs_b).float().mean().item()
print(f"max abs diff = {diff.max().item():.3e}")
print(f"exact mismatch ratio = {mismatch:.4%}")

在 H100 上,bfloat16 下 max abs diff 通常在 1e-3 量级,逐元素不完全相等的比例经常超过 30%。注意这里没有任何随机性:没有 dropout,没有采样,权重和输入完全相同。唯一变化的是矩阵乘法的 batch 维度。

这个差异会一路放大。Transformer 有几十上百层,每层的微小扰动经过残差与归一化传播,最终在 logits 上可能翻转 argmax,也就是同一个 prompt 在负载不同时吐出不同的 token。在多轮采样场景,第一个 token 分歧后整条序列就彻底分道扬镳。

二、三个根因:为什么 GPU 上的浮点是"顺序敏感"的

2.1 浮点加法不满足结合律

这是数学层面的根因。IEEE-754 的 + 不满足结合律:(a+b)+c ≠ a+(b+c)。原因是每次加法都要对阶并舍入到有限精度,不同的分组顺序保留的舍入误差不同。

import numpy as np
a, b, c = np.float32(1e8), np.float32(-1e8), np.float32(1.0)
print(np.float32(a + b + c))   # 1.0
print(np.float32(a + (b + c))) # 0.0  <- 灾难性消去

bfloat16 只有 8 位尾数,这个问题被放大到工程上无法忽略的程度。只要归约(reduction)的顺序或分组方式变了,结果就变。

2.2 Kernel 的并行策略随 batch 形状而变

GEMM 的问题规模 (M, N, K) 一旦变化,cuBLAS/cutlass 的 heuristics 会选不同的 tile 配置、不同的 stage 数,甚至切换是否启用 split-k。

split-k 是这里最凶险的一环:当 K 很大而 M、N 较小(推理 decode 阶段的典型形态)时,为了填满 SM,kernel 会把 K 维切给多个 CTA,每个算部分和,最后用 atomicAdd 累加到输出。

// split-k 的尾部累积(简化示意)
__global__ void gemm_splitk(...) {
    float acc = 0.f;
    for (int k0 = k_start; k0 < k_end; k0 += BK) {
        acc += compute_tile(...);        // 部分和
    }
    // 关键:atomicAdd 的到达顺序由硬件调度决定,每次运行都可能不同
    atomicAdd(&C[row * N + col], acc);
}

atomicAdd 在浮点上是非确定性的:硬件以任意顺序完成这些原子加。于是同一个 (M,N,K) 跑两次都可能得到不同结果——这已经不是 batch 依赖,而是运行间依赖。

2.3 动态 batch 组成(Batch Invariance 的破缺)

现代推理引擎用 Continuous Batching 把不同长度、不同到达时刻的请求拼成一个 batch。序列长度分布一变,kernel 的归约维度就变。于是:

  • 请求 A 在"空闲时段"单独被调度 → 走 batch=1 的路径;
  • 请求 A 在"高峰时段"和 31 个请求合批 → 走 batch=32 的路径;
  • 两者数值不同 → 输出 token 不同。

这就是所谓 batch invariance 破缺:系统的输出依赖于系统中其他无关请求的存在。它让线上问题无法离线复现,是 SRE 和算法工程师共同的噩梦。

三、为什么 2026 年这件事变紧急了

三个真实场景把这个问题从"学术洁癖"推到了"生产阻塞":

1)强化学习后训练(RLVR / GRPO)的 on-policy 偏差。 推理引擎负责 rollout,训练引擎负责前向与反向。如果训练侧算出的 logprob 与推理侧采样时的 logprob 不同,重要性采样比率就不再是 1,策略梯度的假设被破坏,训练曲线会出现莫名的崩溃或 reward hacking。很多"复现不出论文结果"的根源就在这里。

2)评测与灰度 A/B 的可复现。 模型升级后做回归测试,准确率掉了 0.4%——是真退化还是数值噪声?没有确定性基线,这个问题无法回答。

3)合规与审计。 金融、医疗场景需要"同样的输入永远得到同样的输出"作为可审计性前提。一个不可复现的决策链路在审计上是不可接受的。

四、解法一:Batch-Invariant Kernel

核心原则只有一句:让归约顺序只依赖于"数据本身",不依赖于"并行度"和"batch 形状"。

具体到实现,需要同时满足三条:

  • 固定 tile 切分:不根据 M 的大小动态选择 tile;把 M 维的切分固定为常量,超出部分 padding 到固定粒度。
  • 禁用或确定化 split-k:要么完全禁用 split-k(牺牲小 batch 的 SM 占用率),要么改成两阶段确定化归约——先在固定大小的 workspace 中写入每个 k-slice 的部分和,再按固定顺序串行累加。
  • 避免浮点 atomic:把 atomicAdd(float*) 换成"写部分和 + 确定性二次归约"。

用 Triton 表达一个 batch-invariant 的归约内核:

import triton
import triton.language as tl

@triton.jit
def deterministic_reduce(X, Y, N, BLOCK_N: tl.constexpr, SPLIT: tl.constexpr):
    """把归约拆成固定的 SPLIT 段,每段内顺序累积,段间按固定下标累加。
    归约顺序与运行时并行度无关。"""
    pid = tl.program_id(0)
    # 每一段处理固定数量的元素,段的边界只由 N 与常量决定
    seg = N // SPLIT
    acc = tl.zeros([BLOCK_N], dtype=tl.float32)
    for i in range(seg // BLOCK_N):
        offs = pid * seg + i * BLOCK_N + tl.arange(0, BLOCK_N)
        acc += tl.load(X + offs, mask=offs < (pid + 1) * seg, other=0.0).to(tl.float32)
    # 顺序写出:由第二阶段按 0..SPLIT-1 的固定下标归约
    tl.store(Y + pid * BLOCK_N + tl.arange(0, BLOCK_N), acc)

第二阶段是一个纯粹的、固定顺序的串行或树形归约,它对 SPLIT 个部分和按 0..SPLIT-1 的确定顺序累加。这样无论 GPU 当时有多少 SM 空闲、batch 里有多少请求,结果都保持比特级一致。

五、解法二:CUDA 层的确定性原语

不是所有归约都要自己写。CUDA 与主流库已经提供了确定性入口,关键是知道它们在哪、以及默认不开启。

#include <cub/cub.cuh>

// CUB 的确定性归约:对同一输入保证跨运行可复现的结果
void deterministic_sum(const float* d_in, float* d_out, int n) {
    void*  d_temp = nullptr;
    size_t temp_bytes = 0;
    cub::DeviceReduce::Sum(d_temp, temp_bytes, d_in, d_out, n);
    cudaMalloc(&d_temp, temp_bytes);
    cub::DeviceReduce::Sum(d_temp, temp_bytes, d_in, d_out, n);
    cudaFree(d_temp);
}

注意两点:

  1. cub::DeviceReduce::Sum 提供的是 run-to-run determinism(同一 shape 每次结果相同),这正是 2.2 中 atomicAdd 问题的解药;但它不保证 batch invariance——shape 变了归约路径还是可能变。两者要叠加解决。
  2. 对于 GEMM,需要显式关掉启发式对 split-k 的选择。CUBLAS_WORKSPACE_CONFIG=:4096:8(或 :16:8)会限制 cuBLAS 使用额外 workspace,从而禁用部分 split-k 算法,与 torch.use_deterministic_algorithms(True) 配合使用:
import os
os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"   # 必须在 import torch 之前
import torch

torch.use_deterministic_algorithms(True)
torch.backends.cudnn.deterministic = True
# 关键且常被忽略的一行
torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = False

最后一行常被忽略:该选项为 True 时,cuBLAS 会在 bf16 GEMM 中使用低精度累积的中间归约,直接引入额外的顺序敏感误差。确定性场景下必须关掉。

六、解法三:服务端的形状与调度治理

即使 kernel 全部确定化,服务层仍可能引入变异。实践中最有效的四条:

1)固定形状桶(shape bucketing)。 把请求的序列长度 pad 到固定桶(如 512 / 1024 / 2048 / 4096),让 kernel 只面对有限几种 shape。代价是 padding 带来的无效算力,换来的是 shape 空间可枚举、可测、可回归。

2)固定 batch 组合策略。 不要用"当前队列里有多少就组多大"的动态合批,改为固定 batch size 的窗口(如严格 8 / 16 / 32),不足则等待或 padding。这直接消除 batch invariance 破缺。

3)采样确定化。 把采样从"全局 RNG 流 + 多线程竞争"改为"每条请求一个可复现的 counter-based RNG(如 Philox)",由 (request_id, step) 派生:

import torch

def philox_for(request_id: int, step: int, device="cuda"):
    g = torch.Generator(device=device)
    g.manual_seed((request_id * 1000003 + step) & 0xFFFFFFFFFFFFFFFF)
    return g

4)版本号钉死。 记录并锁定推理引擎版本、kernel 库版本、驱动版本、编译 flags。cuBLAS 的一次小版本升级就可能换掉 heuristic,让三个月前的结果彻底无法复现。

七、代价是多少?什么时候不该上

确定性不是免费的。实测区间(H100,Llama-3 8B 级别模型):

手段吞吐损失说明
固定 shape 桶 + padding5%—15%取决于请求长度分布的离散程度
禁用 split-k / 确定化归约8%—20%小 batch decode 阶段代价最大
固定 batch 窗口3%—10%引入排队等待,抬高尾延迟
关闭 bf16 reduced precision reduction2%—5%收益极高,几乎必开

叠加后整体 15%—30%。因此判断标准是清楚的:

  • 必须做确定性:RL 后训练的 rollout 引擎、离线评测基线、需要审计留痕的在线决策、跨版本回归测试。
  • 不必做:纯粹的内容生成、对话、摘要等"用户不比对逐字节结果"的场景。为它们付 20% 吞吐是浪费。

一个现实的折中是双集群:rollout 与评测跑确定性栈(batch-invariant kernel + 固定桶),线上对话跑性能栈,两边共用同一份权重与同一套 tokenizer,只是 kernel 配置不同。

八、生产落地检查清单

  1. 在 CI 里加一条"确定性对拍"用例:同一批 prompt,batch=1 与 batch=32 的输出必须逐 token 相等,否则 fail。这是唯一能防止回归的机制。
  2. 记录每次推理的 engine_version / cublas_version / driver_version / kernel_flags,写入响应 header。
  3. 对 RL 训练链路,额外比对 rollout 侧与 training 侧的 logprob,KL 超阈值就报警——这是 on-policy 破缺的直接信号。
  4. 把 CUBLAS_WORKSPACE_CONFIG 与 use_deterministic_algorithms 做成部署模板的一部分,而不是业务代码里的可选项。
  5. 对自建 kernel,强制 code review 检查点:任何作用于浮点的 atomicAdd、任何按 SM 数量或 occupancy 决定归约顺序的逻辑,都需要显式说明其确定性保证。

九、结语

LLM 推理引擎过去三年的主线是"更快",而 2026 年正在分出第二条主线——"更可复现"。这条线的驱动力不是性能优化,而是LLM 正在从"生成内容的工具"变成"被训练和被审计的系统组件"。一旦推理结果参与梯度更新、参与 A/B 决策、参与合规留痕,不确定性就从可容忍的噪声变成了必须消除的缺陷。

技术上,它要求工程师重新审视一个长期被忽视的假设:浮点运算的顺序是被硬件悄悄决定的。把顺序夺回、把它固定成输入的函数而不是调度器的函数,就是确定性工程的全部要义。这件事不炫,但它决定了你的 RL 能不能训稳、你的评测能不能信、你的线上问题能不能复现。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部