用 torch.cumsum 手写采样逻辑竟然和 torch.
之前在复现 LLM 的 token 采样数值行为时,我掉进了一个挺隐蔽的坑里。原本以为只要逻辑是对的,无论用框架内置函数还是自己写逆变换采样(Inverse Transform Sampling),结果应该是一模一样的,结果在 float32 精度下发现两者居然有偏差。
如果你现在还在纠结怎么让采样结果更稳定,或者在做数值分析,建议直接抛弃手写的累加逻辑。除非你把所有计算都强行提升到 float64,否则你永远无法在数值层面上让
下一篇
用 AI 设计的药真能让人变年轻 6 岁吗 →
我当时的操作逻辑很简单,先拿到 logits 走一遍 softmax 变成概率分布 probs,然后尝试用两种方式取下一个 token。第一种是标准的 PyTorch 做法:
next_token = torch.multinomial(probs, num_samples=1)第二种是我自己写的 CDF 累加法:cdf = torch.cumsum(probs, dim=-1)
u = torch.rand(1, device=probs.device)
next_token = torch.searchsorted(cdf, u)从数学定义上讲,这两者完全等价,都是在做类别分布采样。但实际跑起来后,我发现自己写的这个 FP32 CDF 采样器没法完全复现 torch.multinomial 的行为。
最让我困惑的细节在于:在 LLM 巨大的词表(Vocabulary)面前,float32 的精度其实是很危险的。我观察到在执行 torch.cumsum 之后,会出现一种很诡异的情况,就是 cdf[i] == cdf[i - 1]。理论上只要 probs[i] > 0,累加值应该是增加的,但在有限精度下,如果某个 token 的概率极低,加到那个巨大的累积值上时,由于舍入误差,结果根本没变。这就导致这个 token 对应的区间宽度变成了 0,实际上它在我的手写采样逻辑里被“抹除”了,永远不可能被抽中。
而 torch.multinomial 显然在底层做了优化,它并没有简单地走一遍 cumsum -> rand -> searchsorted 这种线性流程。虽然 PyTorch 的官方文档只描述了它符合什么样的分布,但没说具体怎么实现的。实际上,这种内置函数在 CUDA 和 CPU 上的实现逻辑可能完全不同,而且为了保证数值稳定性,它可能采用了更复杂的算法(比如 Alias Method 或者经过优化的累加策略)来避免这种精度丢失导致的“概率归零”现象。
通过这次踩坑,我有几个比较具体的结论:
- 精度陷阱: 不要试图在 float32 下用
torch.cumsum模拟随机采样,尤其是词表规模在几万甚至十几万的时候,长尾分布的 token 很容易因为精度问题在 CDF 中丢失。 - 实现差异:
torch.multinomial并不是一个简单的包装函数,它在处理浮点数边界和随机数映射时,比手动写searchsorted要健壮得多。 - 复现标准: 如果你的目标是复现一个基于 PyTorch 的 LLM 生成行为,必须直接调用
torch.multinomial(probs, 1),而不是假设“随机采样”就等同于某种特定的逆 CDF 实现。
如果你现在还在纠结怎么让采样结果更稳定,或者在做数值分析,建议直接抛弃手写的累加逻辑。除非你把所有计算都强行提升到 float64,否则你永远无法在数值层面上让
searchsorted 完美对齐框架的内置采样器。 免费 AI 工具箱 · 全部完全免费
更系统的工具评测汇总在AI工具实测笔记,有不少直接可参考的案例。