← 返回文章档案

LM Loss 阅读补充

围绕语言模型损失函数整理相关知识,包括恰当评分规则、梯度与凸性,以及损失函数与激活函数的配套关系。

这篇文章用于沉淀和分析阅读苏剑林《除了交叉熵,LM Loss 还有什么选择?》时涉及的相关知识。内容会随阅读逐步增加。

对损失函数的要求

LLM 每次只看到一个真实 Token,但我们希望它能拟合这个位置所有可能 Token 的完整概率分布 p\boldsymbol{p}

假设我们当前估计的概率分布是 p\boldsymbol{p},我们真正想优化的就是 L(p,q)L(\boldsymbol{p},\boldsymbol{q}),即最小化两个分布间的距离。

核心问题在于,我们没法知道训练语料的真实分布,而只能取样得到样本 ipi \sim \boldsymbol{p} 进行训练。

那么此时就需要保证样本级损失 S(q,i)S(\boldsymbol{q},i) 的期望和整体分布的损失 L(p,q)L(\boldsymbol{p},\boldsymbol{q}) 一致,即

L(p,q)=EipS(q,i)L(\boldsymbol{p},\boldsymbol{q})=\mathbb{E}_{i\sim \boldsymbol{p}}S(\boldsymbol{q},i)

这样才能用样本近似总体损失。

同时,根据损失函数的基本要求,一个用来学习概率分布的损失函数,应该在模型预测分布 q\boldsymbol{q} 等于真实分布 p\boldsymbol{p} 时取得最小值,即:

q=argminqL(p,q)=p\boldsymbol{q}^{*}=\underset{\boldsymbol{q}}{\arg\min} L(\boldsymbol{p},\boldsymbol{q})=\boldsymbol{p}

超平面与超曲面

根据基本要求,定义损失最小值 H(p)L(p,p)=minqL(p,q)H(\boldsymbol{p})\triangleq L(\boldsymbol{p},\boldsymbol{p})=\underset{\boldsymbol{q}}{\min} L(\boldsymbol{p},\boldsymbol{q})

根据定义,L(p,q)=EipS(q,i)=ipiS(q,i)L(\boldsymbol{p},\boldsymbol{q})=\mathbb{E}_{i\sim \boldsymbol{p}}S(\boldsymbol{q},i)=\sum_{i}p_iS(\boldsymbol{q},i),其中 SSp\boldsymbol{p} 无关,故为线性。

那么当固定 q\boldsymbol{q}L(p,q)=ipiS(q,i)=ipiSiL(\boldsymbol{p},\boldsymbol{q})=\sum_{i}p_iS(\boldsymbol{q},i)=\sum_{i}p_iS_i 等价于描述了一个 nn 维的超平面

{(p,L(p,q)):pΔn1}\left\{(\boldsymbol{p},L(\boldsymbol{p},\boldsymbol{q})) : \boldsymbol{p} \in \Delta^{n-1}\right\}

注意,因为 ipi=1\sum_{i}p_i=1,所以它的自由度是 n1n-1,加上结果的一维,共 nn

同时,由于 H(p)H(\boldsymbol{p}) 也是一个标量函数,不过它没有线性约束,故可以看作为 nn 维空间内的超曲面,并且在 p=q\boldsymbol{p}=\boldsymbol{q} 的位置,两平面相切(假设极小值是唯一的)。也就是说:

{(p,L(p,p)):pΔn1}\left\{(\boldsymbol{p},L(\boldsymbol{p},\boldsymbol{p})) : \boldsymbol{p} \in \Delta^{n-1}\right\}

都是 HHp\boldsymbol{p} 处的一张支撑超平面。

因此,HH 的每一个点都存在一张位于其上方的支撑超平面,这正是凹函数的几何特征,所以 HH 是一个凹函数。

用点法式方程描述 H(p)H(\boldsymbol{p})q\boldsymbol{q} 处的切平面:

H(q)+(pq)H(q)=0H(\boldsymbol{q})+(\boldsymbol{p}-\boldsymbol{q})\cdot\nabla H(\boldsymbol{q})=0

改写成线性形式:

p[H(q)+H(q)qH(q)]=0\boldsymbol{p}\cdot[H(\boldsymbol{q})+\nabla H(\boldsymbol{q})-\boldsymbol{q}\cdot\nabla H(\boldsymbol{q})]=0

故所有满足条件的评分函数 S(q,i)S(\boldsymbol{q},i) 有以下形式:

ipiS(q,i)=p[H(q)+H(q)qH(q)]\sum_{i}p_iS(\boldsymbol{q},i)=\boldsymbol{p}\cdot[H(\boldsymbol{q})+\nabla H(\boldsymbol{q})-\boldsymbol{q}\cdot\nabla H(\boldsymbol{q})]

两边逐分量对齐,对于类别 ii,有:

S(q,i)=H(q)+iH(q)qH(q)=H(q)+(eiq)H(q)\begin{aligned} S(\boldsymbol{q},i) &= H(\boldsymbol{q})+\partial_i H(\boldsymbol{q})-\boldsymbol{q}\cdot\nabla H(\boldsymbol{q}) \\ &= H(\boldsymbol{q})+(\boldsymbol{e}_i-\boldsymbol{q})\cdot\nabla H(\boldsymbol{q}) \end{aligned}