LM Loss 阅读补充
围绕语言模型损失函数整理相关知识,包括恰当评分规则、梯度与凸性,以及损失函数与激活函数的配套关系。
ESSAY-003 Hongliang Cao
这篇文章用于沉淀和分析阅读苏剑林《除了交叉熵,LM Loss 还有什么选择?》时涉及的相关知识。内容会随阅读逐步增加。
对损失函数的要求
LLM 每次只看到一个真实 Token,但我们希望它能拟合这个位置所有可能 Token 的完整概率分布 p。
假设我们当前估计的概率分布是 p,我们真正想优化的就是 L(p,q),即最小化两个分布间的距离。
核心问题在于,我们没法知道训练语料的真实分布,而只能取样得到样本 i∼p 进行训练。
那么此时就需要保证样本级损失 S(q,i) 的期望和整体分布的损失 L(p,q) 一致,即
L(p,q)=Ei∼pS(q,i)
这样才能用样本近似总体损失。
同时,根据损失函数的基本要求,一个用来学习概率分布的损失函数,应该在模型预测分布 q 等于真实分布 p 时取得最小值,即:
q∗=qargminL(p,q)=p
超平面与超曲面
根据基本要求,定义损失最小值 H(p)≜L(p,p)=qminL(p,q)
根据定义,L(p,q)=Ei∼pS(q,i)=∑ipiS(q,i),其中 S 和 p 无关,故为线性。
那么当固定 q,L(p,q)=∑ipiS(q,i)=∑ipiSi 等价于描述了一个 n 维的超平面
{(p,L(p,q)):p∈Δn−1}
注意,因为 ∑ipi=1,所以它的自由度是 n−1,加上结果的一维,共 n 维
同时,由于 H(p) 也是一个标量函数,不过它没有线性约束,故可以看作为 n 维空间内的超曲面,并且在 p=q 的位置,两平面相切(假设极小值是唯一的)。也就是说:
{(p,L(p,p)):p∈Δn−1}
都是 H 在 p 处的一张支撑超平面。
因此,H 的每一个点都存在一张位于其上方的支撑超平面,这正是凹函数的几何特征,所以 H 是一个凹函数。
用点法式方程描述 H(p) 在 q 处的切平面:
H(q)+(p−q)⋅∇H(q)=0
改写成线性形式:
p⋅[H(q)+∇H(q)−q⋅∇H(q)]=0
故所有满足条件的评分函数 S(q,i) 有以下形式:
i∑piS(q,i)=p⋅[H(q)+∇H(q)−q⋅∇H(q)]
两边逐分量对齐,对于类别 i,有:
S(q,i)=H(q)+∂iH(q)−q⋅∇H(q)=H(q)+(ei−q)⋅∇H(q)