从一串乱码到崩溃的程序:Softmax计算概率的工程实战
从一串乱码到崩溃的程序:Softmax计算概率的工程实战
当你调用一个大语言模型生成文本时,它在每一步其实并不是直接“说出”下一个词,而是先吐出一串原始浮点数。比如对于词表中的四个候选Token,模型可能输出 [12.5, -3.2, 8.7, 1000.0]。这些数字被称为Logits,它们本身没有概率含义,甚至可以是负数或远超1的值。要将它们转化为“下一个Token是某个词”的概率分布,我们必须使用Softmax函数。
然而,如果你试图用最直观的数学公式来实现这个转换,程序很可能会在特定输入下崩溃。下面这段完全可运行的Python脚本,复现了这个经典陷阱:
import numpy as np
# 模拟大模型输出层的原始Logits(float64),包含一个极端大值
# 注:若使用float32,exp(88)左右即会溢出,此处为演示清晰采用float64
logits = np.array([12.5, -3.2, 8.7, 1000.0])
def naive_softmax(x):
"""朴素实现:直接套用Softmax数学定义"""
exp_x = np.exp(x) # 对每个元素求指数
return exp_x / np.sum(exp_x) # 归一化
result = naive_softmax(logits)
print("计算结果:", result)
运行这段代码,你不会得到期望的概率分布,反而会看到类似 RuntimeWarning: overflow encountered in exp 的警告,输出结果全是 nan 或 inf。问题出在 np.exp(1000.0) 这一步:IEEE 754双精度浮点数能表示的最大有限值约为 1.8e308,而 exp(1000) 的理论值远超此限,导致上溢。更隐蔽的是,即使输入没有大到触发上溢,在FP16等低精度下,较小的负Logits经过exp后也可能下溢为零,使对应Token的概率丢失。