从嵌入表到 logits
目标
只用标准库做出从令牌编号变成向量、再由该向量变回整个词表的分数的全过程。用两种方法计算,确认查表与 one-hot 向量乘以矩阵给出相同的值;数出参数数量;做出在输出一侧把同一张表再用一次的权重绑定;并用数字看清余弦与点积产生分歧的地方。最后测量把嵌入乘以 √d 后只有大小改变、方向不变。
为什么重要
谈论模型时,人们最常出错的地方就在这里。如果把它分开想成“嵌入是查表,输出是矩阵乘法”,就看不出两者是同一张表,也解释不了参数数量为什么对不上。为什么增大词表的提议最后变成内存会议,不数一数 V 乘以 d 也无从体会。
本实验不调用实际模型。这个 Pod 的系统 Python 里没有 numpy、torch、transformers(numpy 只在 /opt/onnx-lab/bin/python 里有)。所以不使用实际模型的词表大小或参数数量之类的数字——只使用在你做出的表上测得的值。
判定中有很多实数比较。评分器用 abs(a-b) <= atol + rtol*abs(b) 来看,同时也看整数判定,如形状、参数数量、近邻列表。所以相加的顺序不同而使最后一位晃动,是可以接受的,真正错误的实现则会被筛掉。
评分器不会相信你写下的说明。它会真正导入你的模块,每次用不同的表和不同的编号调用函数,并与评分器另行计算的值对照。
步骤
- 在 /root/work/tf-embed/embed.py 中创建
VOCAB = 512、DIM = 64、SEED = 20260917以及make_table(vocab, dim, seed)、shape(matrix)、param_count(vocab, dim)。种子相同,得到的表必须永远相同。 - 增加
lookup(table, ids),把编号列表变成向量列表。相同的编号给出相同的向量,词表之外的编号则是IndexError。 - 增加
one_hot(token_id, vocab)、row_times_matrix(vec, matrix)、lookup_via_one_hot(table, ids)、one_hot_mults(vocab, dim, count),看查表与 one-hot 乘法是否给出相同的值。 - 做出
logits(table, hidden),用输入时用的那张表给整个词表打分。不许新建矩阵。 - 用
tied_params(vocab, dim)、untied_params(vocab, dim)、vocab_growth(vocab, dim, factor)数出绑定与分开放时数的个数。 - 做出
dot、norm、cosine、stretch、nearest_by_cosine、nearest_by_dot,看只放大一行时两个列表如何产生分歧。 - 做出
rms(vec)和scaled_lookup(table, ids, dim),测量乘以 √d 后大小变成 √d 倍。 - 用定好的常量全部测量,记录到 /root/work/tf-embed/embed_report.json 和 /root/work/tf-embed/embed_report.md 中。
参考
- 执行契约:评分器会把
/root/work/tf-embed/embed.py当作 Python 模块导入,直接使用VOCAB、DIM、SEED、make_table、shape、param_count、lookup、one_hot、row_times_matrix、lookup_via_one_hot、one_hot_mults、logits、tied_params、untied_params、vocab_growth、dot、norm、cosine、stretch、nearest_by_cosine、nearest_by_dot、rms、scaled_lookup。它不会作为脚本运行,所以可以没有if __name__ == "__main__"。 make_table(vocab, dim, seed)创建一个random.Random(seed),从 0 号令牌的第 0 列开始按行填充。每一列取一次random()并减去0.5。这样评分器才能另行做出同样的表来核对值。shape(matrix)返回(줄 수, 칸 수)(占位符依次为行数与列数),如果各行的列数不同,就抛出ValueError。空表是(0, 0)。lookup(table, ids)对词表之外的编号抛出IndexError。负数也在词表之外——如果放任 Python 的负索引不管,它会从后往前数,悄悄地取出错误的一行。row_times_matrix(vec, matrix)用长度为 V 的行向量和 V×d 的表,得出长度为 d 的向量。即out[c] = sum(vec[r] * matrix[r][c] for r in range(V))。one_hot_mults(vocab, dim, count)是乘法次数。每个令牌是 V 乘以 d 次,用查表则是 0 次。不要测时间,要数这个数。logits(table, hidden)是长度为 V 的列表。out[t] = sum(table[t][c] * hidden[c] for c in range(d)),如果隐藏向量的长度与表的宽度不同,就是ValueError。原样使用这张表,就是权重绑定。vocab_growth(vocab, dim, factor)是带有vocab、bigger_vocab、dim、tied、bigger_tied、untied、bigger_untied、saved这些键的字典。saved是分开放时减去绑定时的值。cosine(a, b)遇到长度为 0 的向量时是0.0。rms([])也是0.0。stretch(table, token_id, factor)返回只把那一行放大factor倍的新表。不要原地修改传进来的表。nearest_by_cosine、nearest_by_dot按分数从大到小返回k个编号。排除自身,分数相同时编号小的在前。scaled_lookup(table, ids, dim)是把lookup的结果乘以math.sqrt(dim)。- 第 8 步用
VOCAB = 512、DIM = 64、SEED = 20260917做出表,并以 137 号令牌作为探测对象。近邻取k = 5,要放大的行是余弦近邻的第五个(从 0 开始数是第 4 个位置),放大的倍数是4.0。增大词表的倍数是2。 - 这个 Pod 没有互联网。
pip install无法使用,在系统 Python 中import numpy也不可用。只用math和random就足够了。 - 官方文档:Attention Is All You Need · PyTorch — MultiheadAttention · NumPy — matmul · Python — math
- 常见错误:按列填充表,导致种子顺序错位;让负编号直接通过;在 one-hot 乘法中把行和列颠倒;在 logits 中另建矩阵来用;没有把分开放时的数按两倍来数;在余弦中没有除以长度;
stretch修改了原表;乘的是 d 而不是 √d。
表的形状与参数数量
在 /root/work/tf-embed/embed.py 中创建 VOCAB = 512、DIM = 64、SEED = 20260917 以及 make_table(vocab, dim, seed)、shape(matrix)、param_count(vocab, dim)。make_table 用一个 random.Random(seed) 从 0 号令牌的第 0 列开始按行填充,每一列从 random() 中减去 0.5。
表是列表的列表。只创建一次 rng = random.Random(seed),每行取 dim 个,就能保持按行的顺序。如果按列去循环,即使种子相同也会得到不同的表。shape 要检查每行的列数,不同就抛出 ValueError——形状错位的表,到后面会悄悄地给出奇怪的值。param_count 是乘积,不是和。
按编号取出行
增加 lookup(table, ids)。把编号列表变成向量列表。同一个编号出现两次,就必须两次得到完全相同的向量,词表之外的编号(包括负数)则是 IndexError。
取出行就是全部。不过 table[-1] 在 Python 里会悄悄地给出倒数第一行,所以必须亲自确认 0 <= token_id < len(table)。取出的行用 list() 复制后再返回,调用方就没有改动原表的风险。同一个编号得到同一个向量,不是 bug 而是性质——嵌入里没有上下文,上下文由后面的注意力来提供。
查表与 one-hot 乘法是相同的值
增加 one_hot(token_id, vocab)、row_times_matrix(vec, matrix)、lookup_via_one_hot(table, ids)、one_hot_mults(vocab, dim, count)。用 one-hot 乘法算出的值必须与 lookup 相同,one_hot_mults 返回每个令牌 V 乘以 d 次的乘法次数。
row_times_matrix 的一格是 sum(vec[r] * matrix[r][c] for r in range(V))。如果把行和列颠倒,长度首先就对不上,所以先用 shape 确认再循环。one-hot 只有一个位置是 1.0,所以乘完再相加,只有那一行留下来——值应该恰好相同才正常。不要想着测时间,要数乘法次数。查表是 0 次。
用同一张表打分
做出 logits(table, hidden)。把一个隐藏向量变成整个词表的分数。必须原样使用输入时用的那张表(权重绑定)。如果隐藏向量的长度与表的宽度不同,就是 ValueError。
对每一行求与隐藏向量的点积,就得到长度为 V 的列表。如果另建矩阵来用,值会完全不同——所谓绑定,就是把那张表再用一次。有一个很好的确认办法:把某个令牌的嵌入原样放进 hidden 试试。与自身的点积是长度的平方,所以该令牌的分数最大。
绑定之后会减少多少
做出 tied_params(vocab, dim)、untied_params(vocab, dim)、vocab_growth(vocab, dim, factor)。vocab_growth 是带有 vocab、bigger_vocab、dim、tied、bigger_tied、untied、bigger_untied、saved 这些键的字典,saved 是分开放时减去绑定时的值。
绑定时表是一份,分开放时同样形状是两份。把词表增大 factor 倍时,宽度明明没动,两个值却都放大了同样的倍数——这就是词表大小的代价。这里全是整数判定,所以不要用实数计算再舍入。
长度改变了顺序
做出 dot(a, b)、norm(vec)、cosine(a, b)、stretch(table, token_id, factor)、nearest_by_cosine(table, token_id, k)、nearest_by_dot(table, token_id, k)。近邻是按分数从大到小的 k 个编号,排除自身。分数相同时,编号小的在前。
cosine 是点积除以两个长度之积,长度为 0 时没有方向,所以是 0.0。stretch 必须创建新表——如果原地修改原表,后面的判定会全部错位。排序用 key=lambda item: (-점수, 번호)(占位符依次为分数与编号)这一行,就连同分数相同时的规则一起写出来了。把一行放大之后,把两个列表并排看。余弦列表不变,而点积列表中被放大的那一行跑到了前面。
乘以 √d,只有大小改变
做出 rms(vec) 和 scaled_lookup(table, ids, dim)。rms 是均方根,空向量是 0.0。scaled_lookup 是把 lookup 的结果乘以 math.sqrt(dim)。
这是给每一格乘同一个数,所以方向丝毫不变——测一下余弦,与乘之前相同。变的只有大小,而且恰好是 math.sqrt(dim) 倍。如果直接乘 dim,大小就变成 d 倍,成了完全不同的值。rms 是把点积除以格数,再开平方根。
把测出来的结果留成记录
用 VOCAB = 512、DIM = 64、SEED = 20260917 做出表,以 137 号令牌为探测对象全部测量。近邻取 k = 5,要放大的行是余弦近邻的第五个,倍数是 4.0,增大词表的倍数是 2。在 /root/work/tf-embed/embed_report.json 中写入 vocab、dim、seed、probe、table_shape、params、tied_params、untied_params、saved、bigger_vocab、bigger_tied、bigger_untied、lookup_mults、one_hot_mults、max_abs_diff、repeat_same、logit_len、logit_argmax、logit_argmax_is_self、cos_neighbors、dot_neighbors、neighbors_differ、stretch_target、stretch_factor、cos_after_stretch、dot_after_stretch、cos_unchanged_by_scale、rms_plain、rms_scaled、rms_ratio,并在 /root/work/tf-embed/embed_report.md 中用 ## 무엇을 쟀나(韩文,意为“测量了什么”)、## 조회와 원-핫 곱은 같은 연산이다(韩文,意为“查表与 one-hot 乘法是同一种运算”)、## 가중치를 묶으면 무엇이 줄어드나(韩文,意为“绑定权重会减少什么”)、## 내적과 코사인이 갈리는 자리(韩文,意为“点积与余弦产生分歧的地方”)、## √d 를 곱하면 무엇이 달라지나(韩文,意为“乘以 √d 会有什么变化”)五节来写。
数字不要手写,要用实际运行你的代码得到的值来填。one_hot_mults 是查137 号令牌一个时的值,lookup_mults 是 0。max_abs_diff 是 lookup 与 lookup_via_one_hot 的值之差中最大的绝对值。repeat_same 表示同一个编号查两次时两个向量是否相同。logit_argmax 是把 137 号的嵌入原样放进隐藏向量时分数最大的编号。cos_unchanged_by_scale 表示乘以 √d 之前和之后的余弦是否相同。rms_ratio 是乘之后的大小除以乘之前的大小得到的值。