编号变成向量,向量再变成分数
一句话总结
嵌入只是一张取出某一行的表,而输出一侧的 logits 则是把这张表再用一次的乘法。两者之间包含的数的个数是词表大小乘以宽度,增大词表的代价就在这里付出。
为什么需要它
分词做完之后,文本就成了整数列表。可是整数不能直接使用。编号为 3 的令牌和编号为 4 的令牌并不是邻居,但只要用数字表示,所有运算都会这样解读。所以给每个编号配一个向量。把这些向量堆叠成一张,就是嵌入表。
这里第一个让人卡住的是为什么是查表。教科书写的是“用 one-hot 向量乘以矩阵”,而代码却是 table[token_id] 一行。两者看起来像是不同的事情。其实并无不同——one-hot 只有一个位置是 1,其余都是 0,乘完再相加,只有那一行留了下来。值丝毫不差,只是乘法次数不同。查表不是另一种运算,而是同一种运算的捷径。
第二个让人卡住的是出口。走完注意力和各个块之后,剩下一个向量。要把它重新变回词,这是给整个词表打分的事。要把宽度为 d 的向量扩展成大小为 V 的分数,需要一个 V×d 的矩阵——而这种形状的矩阵其实已经有了。输入一侧的嵌入表就是这个形状。
查表只是乘法的捷径
表的形状是(词表大小 V,宽度 d)。它包含的数的个数是 V 乘以 d。宽度不动,只把词表翻一倍,这个个数也翻一倍。为了减少令牌数而增大了词表,但代价要从这张表里出。
ids = [7, 7, 41]
rows = [table[i] for i in ids] # 조회
# 같은 값을 원-핫으로 계산하면
one = [0.0] * V; one[7] = 1.0
row = [sum(one[r] * table[r][c] for r in range(V)) for c in range(d)]
# rows[0] 과 row 는 같은 값이다. 곱셈만 V 곱하기 d 번 더 했다.
这里还能看出一点。ids 的前两个位置是同一个编号,所以得到的是完全相同的向量。前面有什么、后面来什么,都一样。嵌入里没有上下文。同一个词因位置不同而被读成不同含义,这是注意力在后面做的事,表只是表。torch.nn.MultiheadAttention 的第一个参数是 embed_dim,也是这个缘故——嵌入的宽度就是模型的宽度,表里定下的 d 会被后面所有层原样接收。
在实际使用矩阵乘法的地方,会由 NumPy 的 matmul 之类的来代为计算。不过本实验 Pod 的系统 Python 里没有 numpy,只在 /opt/onnx-lab/bin/python 里有,所以这里用标准库亲自计算两种方法,看值是否相同。
出口:把同一张表再用一次
Attention Is All You Need 在讲嵌入的一个短小章节里写了两件事。一是两个嵌入层和 softmax 之前的线性变换共享同一个权重矩阵,二是在嵌入层把这个权重乘以 √d。前一件就是权重绑定(weight tying)。
绑定之后会出现两件事。第一,表只有一份,所以数的个数减半。分开放的话,V·d 有两份,是 2·V·d。第二,打分的方式变成了点积。对于隐藏向量 h,令牌 t 的分数就是表的第 t 行与 h 的点积。所以当 h 与某个令牌的嵌入相同时,该令牌的分数最大——因为与自身的点积是长度的平方,比其他任何点积都容易更大。
这个性质很方便,但陷阱也出在同一个地方。绑定之后,同一张表要同时承担进来的含义和出去的分数两件事。对一侧好的布局,不一定对另一侧也好。所以绑不绑并不是免费的选择,而是一笔交易——参数减半的代价,是让表同时做两件事。
不要用点积去找相近的令牌
寻找“与这个令牌相近的令牌”时,不能直接用点积。因为点积里乘进了对方的长度。即使方向稍微不那么吻合,只要够长,也会排到前面。
余弦把这个长度除掉,把它抹去,只留下方向。确认差别最可靠的办法,是把表中的一行放大若干倍。这一行的余弦丝毫不变(方向没变),点积则全部放大同样的倍数。用余弦取出近邻列表,顺序不变;而用点积取出的话,被放大的那一行会跳到最前面。
所以 logits 的顺序和“含义相近的顺序”不是一回事。logits 是点积的顺序,里面混有长度。
论文乘以 √d 的地方
同一章节的另一句话,是把嵌入乘以 √d。乘了之后有什么变化呢?方向丝毫不变。每一格都乘了同一个数,所以余弦不变。变的只有大小,而且大小恰好变成 √d 倍。
大小为什么重要?在给嵌入加上位置信息的地方,如果两种信号的大小相差太多,其中一个就会被淹没。所以需要事先把大小调齐。本实验不去证明“为什么偏偏是 √d”,而是用数字确认乘了之后大小变成 √d 倍,方向不变。能测出来、能说出口的,到这里为止。
在现场相遇的样子
第一,增大词表的提议,最后变成了内存会议。觉得令牌变少挺好,就提出把词表翻一倍,得到的回答是表也要翻一倍。宽度明明没动,却是这样。哪边划算,要把两个值并排数一数才知道。
第二,因为“绑了还是没绑”,参数数量对不上。明明写下的是同样的配置,算出来的值却差了一张表。这是整整一张表有或没有,所以不是舍入误差之类的问题。
第三,相似令牌列表很奇怪。用点积取出来,却称之为“含义相近”,长度大的行就会插进任何一个 query 里。换成余弦,那一行就消失了。
第四,只取出嵌入,却期待上下文。同一个词在表里永远是同一行。如果需要表示句子的向量,就必须让它通过模型,查表得到的值是没有上下文的值。
第五,改变宽度,后面全都跟着动。嵌入的 d 是后面所有层接收的宽度,不可能只改一处。
实际工作中真正重要的事
- 先数表的形状。把 V 和 d 两个数相乘,就得到这张表占用的空间。讨论增大词表,要先把这个乘积写下来。
- 先确认有没有绑定。参数数量对不上时,这是最先要看的地方。
- 近似用余弦,打分用点积。两者混用,长度会悄悄地改变顺序。
- 要知道查表与乘法是同一种运算。这样就不会害怕优化,反过来也不会说出“因为是查表所以不同”这种错误的解释。
下一项实验要做什么
把 /root/work/tf-embed/embed.py 一步一步做大。只使用标准库——这个 Pod 的系统 Python 里没有 numpy、torch、transformers,numpy 只在 /opt/onnx-lab/bin/python 里有。因为不调用实际模型,所以不使用实际模型的词表大小或参数数量之类的数字。这里出现的值,全部是在你做出的表上测出来的。
从做出确定性的嵌入表开始,数出形状和参数数量,按编号取出行,再用 one-hot 乘法重算同样的值,看两个值是否相同。接着用同一张表算出 logits,数出绑定与分开放时数的个数,并把词表翻一倍试试。
最后两步是要点。把表中的一行放大若干倍,分别用余弦和点积取出近邻列表。余弦列表不变,而点积列表中被放大的那一行会跑到最前面。最后把嵌入乘以 √d,并排测量大小恰好变成 √d 倍、方向不变这两件事。评分器会真正导入你的模块,每次用不同的表和不同的编号调用函数,并与它自己另行计算的值对照。