亲手算出一个损失
目标
从 logits 出发,只用标准库亲手写出把“这个模型做得有多好”变成一个数字的过程。一路做到稳定的 log-softmax、一个位置的损失、多个位置的平均即交叉熵、用指数换回去的困惑度,然后把标签错开一位和去掉填充位置这两件事对这个数字造成的变化并排放在一起测量。最后换成底为 2 的对数,也以比特/令牌来读。
为什么重要
训练也好,评估也好,都是盯着这一个数字来行动的。可是,生成这个数字的过程中,有四处不报错却悄悄出错的地方。对数是什么时候取的,标签有没有错开一位,填充有没有去掉,对数的底是什么。四件事在代码里都只有一行,错了也不会抛异常,而且通常是朝数字变好的方向出错。所以没有起疑的契机。 本实验不调用实际模型。这个 Pod 的系统 Python 里没有 numpy、torch、transformers。取而代之的是,为每个位置确定性地做好针对整个词表的一行分数,在这之上亲手写出同样的计算。所以“某个模型的困惑度是多少”这样的话,这里不说。出现的数字,全部是用你做的数据测出来的。 如果说相邻的模块讲的是从分布中挑出一个的方法(温度、top-k、top-p),那么这里讲的就是测量这个分布错得有多厉害的方法。挑选之前,先要测量。 评分器不会相信你写下的说明。它会真正导入你的模块,每次用不同的 logits 直接调用函数,并与评分器另行计算的值对照。输入每次运行都会变,所以无法把值背下来填进去。
步骤
- 在 /root/work/tf-loss/loss.py 中创建
VOCAB、PAD_ID、SEQ、dataset()以及log_softmax(xs)。不经过概率,直接从 logits 走到对数概率。 - 增加
NEG_INF和naive_log_softmax(xs),做出故意用错顺序的版本。先求概率再取对数,重现在最底部出现-inf的现象。 - 增加
token_loss(logits, target),测量一个位置的损失。它是正确答案令牌的对数概率取负后的值。 - 增加
cross_entropy(rows, targets),对多个位置的损失求平均。 - 做出
perplexity(rows, targets)和uniform_perplexity(vocab_size, length)。确认在均匀分布下困惑度等于词表大小。 - 做出
shift_pairs(rows, ids)、shifted_loss(rows, ids)、unshifted_loss(rows, ids),把标签错开一位和没错开的情形并排测量。 - 做出
kept_positions(targets, pad_id)和masked_cross_entropy(rows, targets, pad_id),去掉填充位置再测量。分子和分母两边都要去掉。 - 增加
bits_per_token(loss)、nats_per_token(bits),并把测得的值记录到 /root/work/tf-loss/loss_report.json 和 /root/work/tf-loss/loss_report.md 中。
参考
- 执行契约:评分器会把
/root/work/tf-loss/loss.py当作 Python 模块导入,直接使用VOCAB、PAD_ID、SEQ、dataset、log_softmax、NEG_INF、naive_log_softmax、token_loss、cross_entropy、perplexity、uniform_perplexity、shift_pairs、shifted_loss、unshifted_loss、kept_positions、masked_cross_entropy、bits_per_token、nats_per_token。它不会作为脚本运行,所以可以没有if __name__ == "__main__"。 VOCAB在 12 以上,SEQ是 16 个以上的令牌编号列表,所有编号必须不小于 0 且小于VOCAB。最后 3 个以上是PAD_ID,在它前面不能有PAD_ID。PAD_ID同样不小于 0 且小于VOCAB。dataset()返回(로짓 줄 목록, 토큰 번호 목록)(占位符依次为 logits 行列表与令牌编号列表)。行数与len(SEQ)相同,每一行的长度是VOCAB。不要使用随机数——调用两次必须得到同样的值。- 请把
dataset()中位置 t 的那一行,做成让位置 t+1 上实际出现的令牌得到最高分。这样标签为什么需要错位,才会在数字上显现出来。句子结束之后的填充位置,要给更大的分数——这是为了说明,把容易猜的位置放进平均,值就会看着变好。 log_softmax(xs)是x_i - (max + log sum exp(x - max))。不能先生成概率、相除之后再取对数。log_softmax([0.0, -800.0])的两格都必须是有限值,第二格在 -800 附近。naive_log_softmax(xs)相反,先生成概率再取对数。math.log(0.0)会抛出异常,所以概率为 0.0 的位置要亲自填成NEG_INF。同样的输入下log_softmax是有限的,而只有这边会得到-inf,这就是这一步的要点。token_loss(logits, target)是-log_softmax(logits)[target]。别忘了把符号翻转。cross_entropy(rows, targets)是平均而不是总和。如果targets为空,就返回 0.0。perplexity(rows, targets)是exp(평균 손실)(韩文,意为“exp(平均损失)”)。不是每个位置各取exp再求平均——在均匀分布下两个值碰巧相同,所以仅凭这一点区分不出来。uniform_perplexity(vocab_size, length)做出length行所有分数都相同的 logits,并测量困惑度。结果必须与vocab_size相同。shift_pairs(rows, ids)是(rows[:-1], ids[1:])。logits 丢掉后面一个,令牌丢掉前面一个。shifted_loss是用这一对测得的交叉熵,unshifted_loss是不错位、直接把rows和ids传进去测得的值。在你的数据上,错位的一边必须更小。masked_cross_entropy(rows, targets, pad_id)只把targets[i] != pad_id的位置相加,并除以这些位置的个数。如果除以整个长度,值就会悄悄变小。如果没有剩下的位置,就是 0.0。- 第 8 步的报告用调用一次
dataset()得到的那一套来测量。在用shift_pairs错位后的一对上测masked_loss,没错位的值用unshifted_loss来测。bits_per_token以masked_loss为基准得出。probe_gap固定为 800,naive_is_inf表示naive_log_softmax([0.0, -800.0])[1]是否为-inf,stable_logprob是log_softmax([0.0, -800.0])[1]。 - 这个 Pod 没有互联网。
pip install无法使用,系统 Python 里没有 numpy、torch、transformers。numpy 只在/opt/onnx-lab/bin/python里有。只要import math就足够了。 - 官方文档:Attention Is All You Need · Python — math · Python — statistics · Hugging Face — Text generation
- 常见错误:先生成概率再取对数;没有翻转损失的符号;用总和代替平均;每个位置各取
exp再求平均;把标签朝相反方向错位;屏蔽时没有改分母;把底为 2 的对数和自然对数混着写。
从 logits 直接走到对数概率
在 /root/work/tf-loss/loss.py 中创建 VOCAB(12 以上)、PAD_ID、SEQ(16 个以上,最后 3 个以上是 PAD_ID)、dataset() 以及 log_softmax(xs)。dataset() 返回 (로짓 줄 목록, 토큰 번호 목록)(占位符依次为 logits 行列表与令牌编号列表),不使用随机数。log_softmax 不经过概率,直接从 logits 得到对数概率。
mkdir -p /root/work/tf-loss。公式是 x_i - (max + log sum exp(x - max)) 一行。没有除法这一点就是要点——先生成概率、相除之后再取对数的话,极小的概率会变成 0.0,对数就崩了。如果 log_softmax([0.0, -800.0]) 的第二格是有限的值(-800 附近),就对了。dataset() 要做成让位置 t 的那一行给位置 t+1 的令牌最高分,在填充是正确答案的位置,则要给更大的分数。
故意把它弄崩看看
增加 NEG_INF 和 naive_log_softmax(xs)。这次先求概率再取对数。概率塌缩成 0.0 的位置,math.log 会抛出异常,所以亲自填成 NEG_INF。请确认同样的输入下 log_softmax 是有限的,而只有这边会得到 -inf。
NEG_INF = float("-inf")。只要换一下顺序——把 exp 后的值除以总和得到概率,再对这个概率取对数。在中间的那些值上会得到与上一步的函数相同的答案,只在最底部产生分歧。请放进像 [0.0, -800.0] 这样差距很大的一行试试。要先看概率是不是 0.0 并过滤掉,如果直接调用 math.log,就会以异常告终。
一个位置的损失
增加 token_loss(logits, target)。它是模型给正确答案令牌的对数概率取负后的值。如果给了正确答案概率 1,就是 0,概率越小,值越大。
一行——-log_softmax(logits)[target]。如果忘了翻转符号,值就全是负数,“损失在下降”这句话就颠倒了。不能直接使用 logits 本身。给其他令牌分配了什么,并不单独去数——因为总和是 1,正确答案的份额就是其余的份额。
多个位置的平均
增加 cross_entropy(rows, targets)。对每个位置求出 token_loss,返回平均。如果 targets 为空,就是 0.0。
是平均而不是总和。如果用总和来测,长句子永远是坏句子,长度不同的文字就无法比较。用 zip(rows, targets) 配对相加,再除以 len(targets)。这里分母里放什么,在第 7 步会再次成为问题。
困惑度的刻度
做出 perplexity(rows, targets) 和 uniform_perplexity(vocab_size, length)。前者是 exp(평균 손실)(韩文,意为“exp(平均损失)”),后者做出所有分数都相同的 logits 行并测量困惑度。请确认结果是否与 vocab_size 相同。
math.exp(cross_entropy(rows, targets)) 一行。容易与“每个位置各取 exp 再求平均”混淆,而在均匀分布下两个值碰巧相同,所以这个测试区分不出来。uniform_perplexity 做出 [[0.0] * vocab_size] * length 这种形式的行,随便放个正确答案就行——所有概率都是 1/V,所以损失是 log V,取指数就是 V。这个值就是困惑度刻度的起点。
标签错开了一位
做出 shift_pairs(rows, ids)、shifted_loss(rows, ids)、unshifted_loss(rows, ids)。shift_pairs 返回 (rows[:-1], ids[1:]),把位置 t 的 logits 与位置 t+1 的令牌配成一对来猜。另外两个是错位版本和没错位版本的交叉熵。
最后一行没有要猜的下一个令牌,而第一个令牌前面没有预测它的那一行。所以 logits 丢掉后面一个,令牌丢掉前面一个。如果朝相反方向错位(rows[1:]、ids[:-1]),就等于把已经看到的东西当作答案交出去,数字会奇怪地变好。在你的数据上,错位一边的损失必须明显小于没错位的一边——因为 dataset() 把分数加在了位置 t+1 的令牌上。
去掉填充再测量
做出 kept_positions(targets, pad_id) 和 masked_cross_entropy(rows, targets, pad_id)。前者是正确答案不是填充的位置的编号列表,后者是只把这些位置相加,再除以这些位置的个数得到的平均。如果没有剩下的位置,就是 0.0。
要去掉的地方有两处——加的一侧和除的一侧。如果漏掉除的一侧,除以整个长度,值就会悄悄变小,而且方向总是朝变好的一边,没有起疑的契机。把 kept_positions 单独拿出来,就能亲眼看到剩下了什么。把这个函数用在错位后的一对上,结果必须比带着填充测出的值大——因为填充是容易猜的位置。
把四个数字并排放在一起
增加 bits_per_token(loss) 和 nats_per_token(bits),并在 /root/work/tf-loss/loss_report.json 中写入 vocab_size、pad_id、seq_len、pad_count、kept、dropped、unshifted_loss、shifted_loss、masked_loss、unshifted_perplexity、shifted_perplexity、masked_perplexity、bits_per_token、uniform_perplexity、probe_gap、naive_is_inf、stable_logprob,并在 /root/work/tf-loss/loss_report.md 中用 ## 무엇을 쟀나(韩文,意为“测量了什么”)、## 한 칸 어긋나면(韩文,意为“错开一位会怎样”)、## 패딩을 빼면(韩文,意为“去掉填充会怎样”)、## 비트로 재면(韩文,意为“用比特来测”)四节来写。
数字不要手写,要用实际运行你的代码得到的值来填。masked_loss 是对用 shift_pairs 错位后的一对使用 masked_cross_entropy 得到的值,shifted_loss 是对同一对不做屏蔽测得的值。bits_per_token 以 masked_loss 为基准得出——请亲自确认 2 ** 그 값(韩文,意为“2 的该值次幂”)是否与 masked_perplexity 相同。uniform_perplexity 是 uniform_perplexity(VOCAB, 8)。probe_gap 固定为 800,naive_is_inf 表示 naive_log_softmax([0.0, -800.0])[1] 是否为 -inf,stable_logprob 是同样的输入下 log_softmax 给出的第二格。