用 torch.cumsum 手写采样逻辑竟然和 torch.

老张在路上 中级 41分钟前 561 浏览 4 点赞 约 2 分钟

之前在复现 LLM 的 token 采样数值行为时,我掉进了一个挺隐蔽的坑里。原本以为只要逻辑是对的,无论用框架内置函数还是自己写逆变换采样(Inverse Transform Sampling),结果应该是一模一样的,结果在 float32 精度下发现两者居然有偏差。

我当时的操作逻辑很简单,先拿到 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 完美对齐框架的内置采样器。
求助pytorchtorch.multinomialfloat32
更系统的工具评测汇总在AI工具实测笔记,有不少直接可参考的案例。

全部回复 (4)

大鹏的日常 初级 32分钟前
记得试下torch.topk先截断一下,能少很多噪声。
0 回复
数据分析师小美 初级 32分钟前
那你最后是用float64解决的,还是直接调阈值对齐的?
0 回复
躺平产品经理 初级 28分钟前
@数据分析师小美 我直接调阈值了,感觉这样最快,你当时怎么处理的?
0 回复
脚本小子阿杰 专家 32分钟前
我也被坑过,fp32累加确实容易飘,建议直接用内置的。
0 回复

发表回复

支持 Markdown 格式