TT Lab
开始
学习 学习路径 课程

Transformer — 手算一遍注意力

把模型的好坏变成一个数字

在 TT Lab 中继续学习

一句话总结

模型给正确答案令牌的概率,取对数、取负之后求平均,就是交叉熵,再把它用指数换回去,就是困惑度(perplexity)。困惑度可以读成“在每个位置平均要在几条岔路之间犹豫”。

为什么需要它

训练也好,评估也好,最终都是盯着一个数字来行动的。损失在下降,就继续跑;不下降,就改点什么。可是,生成这个数字的过程中,有好几处不报错却悄悄出错的地方。

实际遇到的情形是这样的。损失在某个时刻变成了 inf,从下一个批次起 nan 蔓延开来。或者损失曲线好好地在下降,生成的文字却一塌糊涂。或者昨天测的困惑度和今天测的值不一样,而模型明明没变。或者别人论文里的数字和我的数字差了一倍,却不知道哪边错了。

原因通常是这四个中的一个:对数是什么时候取的、标签有没有错开一位、填充位置有没有去掉、对数的底是什么。这四件事在代码里都只有一行,错了也不会抛异常。所以不亲手测一测就看不出来。

什么时候取对数

softmax 把一行分数变成总和为 1 的分布。我们需要的不是那些概率,而是对数概率。这样就会想:“先求出概率再取对数不就行了吗?”,可是这个顺序在最底部就崩了。

双精度实数能容纳的最小正数在 5e-324 附近。如果正确答案令牌的分数比其他的低很多,它的概率就会降到这个数以下,塌缩成 0.0。0 没有对数。Python 的 math.log 会在那里抛出异常,而为了避开异常填进去 -inf 的话,那个位置的损失就成了 +inf,一求平均,整个句子都被染成 inf。

答案是干脆不去生成概率。

# log p_i = x_i - (max + log sum exp(x - max))
top = max(xs)
lse = top + math.log(sum(math.exp(x - top) for x in xs))
logp = [x - lse for x in xs]

这里没有除法。减去大值再 exp,所以不会溢出;对数概率是用减法得到的,所以也不会触底。不论概率有多小,它的对数都只是一个较小的负数而已。框架把 softmax 和损失捆成一个函数来卖,就是这个原因——把两个运算分开调用,信息就会在两者之间丢失。

之前的实验里处理过的 softmax 稳定化,堵住的是上面(exp 溢出)。这里堵的是下面。同样是减去 max,做的是两件事,但崩溃的位置和症状都不同。

从一个位置到整个句子

一个位置的损失是一行式子。

损失(t) = -log p(正确答案令牌 t)

如果给了正确答案概率 1,损失就是 0,概率越小,损失越大。给其他令牌分配了什么,并不单独去数。因为总和是 1,正确答案的份额就是其余的份额。

句子的损失是这些值的平均。是平均而不是总和,是为了比较长度不同的句子。如果用总和来测,长句子永远是坏句子。它做的事与 statistics 里的平均相同,但分母里放什么,到后面会成为问题。

困惑度只是换了刻度

损失 1.06 这个值让人没有感觉。用指数换回去,就成了可以读懂的数字。

困惑度 = exp(平均损失)

有一个确定刻度的办法。放进一个什么都不知道的模型试试。给整个词表同样的分数,概率就是 1/V,损失是 log V,所以困惑度恰好是 V,也就是词表大小。所以如果困惑度在词表大小附近,说明这个模型什么也没学到;如果比这还大,就比均匀分布还差。

由此可以马上得出一个结论。词表不同的两个模型,困惑度不能比较。因为刻度的起点不同。分词器不同,把同一段文字切成的片数也不同,连分母也不一样了。把论文里的数字与我的数字并排放之前,要先看词表和分词器是否相同,原因就在这里。

标签错开了一位

语言模型用位置 t 的输出去猜位置 t+1 的令牌。Attention Is All You Need 的解码器做的就是这件事,Hugging Face 的生成文档所讲解的下一个令牌预测,说的也是同一回事。

所以测损失时,logits 要丢掉后面一个,令牌要丢掉前面一个。因为最后一行没有要猜的下一个令牌,而第一个令牌前面没有预测它的那一行。

如果不错开一位会怎样?不会报错。这等于让模型“去猜你现在正在看的令牌”,只会让损失变差。反过来,该错位的地方错了两次,或者朝相反方向错位,数字反而会奇怪地变好——因为这等于把自己已经看到的东西当作答案交出去。如果损失曲线看着挺像样,生成结果却很糟,就先来看这里。

填充是白送的分数

为了合并成批次,要把短的行补齐到相同长度。补进去的填充令牌不是内容,而是占位。

问题在于填充令牌太容易猜了。句子结束之后总是出现同样的东西,所以模型很快就有了把握。把这些位置放进平均,损失就会降低,填充越多的批次,降得越低。模型明明没变,只改变批次的构成,数字就变好了。

要去掉的地方有两处。加的那一侧要去掉,除的那一侧也要去掉。如果漏掉除的那一侧,就会把剩下位置的损失除以整个长度,值悄悄变小。这个错误尤其难找——因为方向总是朝“变好”的一边,没有起疑的契机。

用比特来测

把对数的底换成 2,单位就变成比特/令牌。只是做一次除法。

比特/令牌 = 自然对数损失 / ln 2

因为只换了底,所以 2 ** 비트(韩文,意为“2 的比特数次幂”)与 exp(자연로그 손실)(韩文,意为“exp(自然对数损失)”)是同样的值。明明只是用不同的尺子读同一个东西,但压缩方面的文献常用比特来写,深度学习方面常用自然对数来写,所以光看数字,会觉得差了将近一倍。抄别人的表之前,必须先确认底。

在现场相遇的样子

第一,损失突然变成 inf 或 nan。这是先生成概率再取对数的代码,碰到了极低概率的瞬间。把它改成用减法得到对数概率,就消失了。

第二,损失在下降,生成的东西却很差。先看标签移位是不是错位了。提前把答案给出去再让它去猜,损失想降多低都能降。

第三,同一个模型的困惑度每次运行都不同。很可能是评估批次的填充比例变了。如果屏蔽做得对,批次构成变了,值也不会晃动。

第四,和别人的数字差了一倍。依次对一对:对数的底、词表大小、分词器、分母是按令牌算还是按词算。通常就是其中之一。

第五,只看损失就上线,结果出了事故。困惑度说的是“下一个令牌猜得有多准”,而不是“给出的回答有没有用”。用来确认它在下降是好的,但只凭它不能说这是个好模型。

实际工作中真正重要的事

下一项实验要做什么

把 /root/work/tf-loss/loss.py 一步一步做大。不是调用实际模型,而是只用标准库亲手做出同样的计算——这个 Pod 的系统 Python 里没有 numpy、torch、transformers,numpy 只在 /opt/onnx-lab/bin/python 里有。所以这里出现的数字,全部是用你做的数据测出来的。

从稳定的 log-softmax 开始。然后另外做一个故意用错顺序的版本,亲眼看看先求概率时 -inf 是否真的会出现。两个函数并排放着,才看得出在哪个位置产生分歧。

在此基础上,一路上升到一个位置的损失、多个位置的平均、用指数换回去的困惑度。还要亲自确认,什么都不知道的模型的困惑度等于词表大小。

最后三步是本实验的要点。对同一份数据,把标签错开与没错开、去掉填充与没去掉填充的数字并排放在一起。四个数字都没有报错,四个看起来都像那么回事,却互不相同——这就是本模块想展示的全部。最后换成底为 2 的对数,也以比特/令牌来读。评分器会真正导入你的模块,每次用不同的 logits 直接调用函数,并与它自己另行计算的值对照。