- 在 Transformer 的 encoder 和 decoder 中对于序列之间的相关性的表达是通过 Attention 机制实现的
- Attention 机制强在可以将当前的选中位置的序列信息与整个序列中的信息构建起相关性,但是这份相关性是不关注序列的语序的,相关性的处理操作是平行进行,或者说并行的
- 但是在传统的基于 RNN 的序列模型(包括升级版的 LSTM 和 GRU)中,输入的序列顺序是模型的内置特性,这些模型能自然地处理输入之间的时间关系
- Transformer 中接收的是具备时序关系的序列信息,但是只依靠Attention机制无法体现时序的概念
- 因此需要一个特殊的概念来表示序列中顺序的概念
- 假设我们有一个句子,"我爱吃苹果"
- 我们可以使用词嵌入(word embedding)来表示每个词,但是这样的表示方式并不能捕捉到句子中词语的顺序
- 如果我们不使用Positional Encoding,那么Transformer模型可能就无法区分 "我爱吃苹果" 和 "苹果爱我吃",
- 因为在它看来,这两个句子的词嵌入是一样的,attention机制在进行相关性计算的是不会采纳其中的顺序信息
- 这时候就需要有一个新的机制能够让 Transformer 模型识别每一个序列信息在序列中的位置,从而规避上述的问题
- 假设现在没有 Positional Encoding 处理,仅有 word Embedding 处理,输入序列将会转换成一个词向量表示
- 提示:为了简化理解,这里使用单头Attention 代替 Transformer 中的多头 Attention
- 这份词向量表示在进入单头Attention机制后,需要与随机初始化的可训练参数矩阵 Query, Key, Value 矩阵进行相乘操作分别得到 q,k,v 向量
- 使用选定的q 与全序列的k 进行相乘得到 original attention score,这个得分是当前选定的序列内容与全序列中其他内容之间的相关性得分
- 将上一步的得分进行压缩和softmax处理后得到 attention 占比
- 使用 attention 占比与 v 进行相乘再相加得到最终的对应序列位置的输出向量
- 在整个过程中,我们发现替换序列中的顺序对attention机制计算得分和计算输出向量的过程没有任何影响
- 换句话来说,attention机制也就无法解决位置相关性的问题
- 了解完了PE对于Transformer模型的必要性后,我们就需要了解如何构建一个 PE
- PE的核心目标是表达出序列之间的顺序性,序列的顺序包含两个数学意义上的概念:
- 一个是序列中的个体与个体之间的前后顺
- 一个是序列中的个体在序列中的顺序
- Positional Encoding在Transformer模型中是一个固定的,不参与训练过程的组成部分
- 这意味着,无论训练过程如何进行,Positional Encoding的值都不会改变
- 因此,Positional Encoding所包含的信息必须足够丰富,能够在不受训练影响的情况下,有效地传达位置信息
- 如果我们使用简单的index encoding,例如[1,2,3,...]等作为位置信息,首先是这些数值的表达能力有限
- 这种表达方式虽然可以告诉模型每个单词的绝对位置,但是它无法有效地表达单词之间的相对位置关系
- 例如,模型很难从index encoding中推断出“位置5和位置10之间的距离是5这样的相对位置概念,因为这些信息无法被函数化”
- 另外,简单的index encoding没有考虑到输入的维度
- 在Transformer模型中,词嵌入通常有很高的维度(例如512或者768)
- 如果我们使用单一的索引值作为位置编码,那么这种编码将很难和高维的词嵌入有效地交互,从而影响模型的学习效果
- 使用简单的index encoding可能无法很好地处理序列长度超出训练集的情况,如果你在训练时没有看到超过某一长度的序列,那么在推理(或测试)时,模型可能无法正确处理超过该长度的新序列
- 举例来说:训练时只有1024长度的max length,但是预测时出现了1248长度的序列,那么index encoding 中的从1025-1248位置的 pe encoding 值是模型没有学习到的
- 这是因为在训练时,模型从未见过这些新的位置编码,因此,可能无法正确解释这些编码
- PE的设计带着两个目标:个体与个体之间的相对位置关系和个体和整体之间的绝对位置关系
- 这两个目标都需要使用函数的方式进行体现和表达
- 我们的研究学者非常聪明,聪明到可以凭空设想到可以使用正弦函数和余弦函数的周期性来进行无限长度的序列的表达
为什么说正弦函数和余弦函数的周期性可以用于表达无限长度序列的唯一性呢,周期性和唯一性是不是存在概念上的冲突呢?
- 首先我们了解,每个位置的编码都是一个高维向量,向量的每一个维度都由一个正弦或余弦函数生成
- 重要的是,每一个维度的函数都有不同的频率,这意味着向量中的每一元素(即每一个维度)都有不同的周期
- 重点:每一个位置都有一个高维向量表示PE,同时高维向量中每一个维度,也就是向量中的每一个元素都具备自己的周期性
- 从整个序列的角度看和单个序列个体的角度区别看,我们追求的是序列整体具备唯一性,也就是序列个体与个体之间的PE表达只要维持住唯一性即可
- 这样,尽管个体的单个维度中的位置编码可能在一个周期后重复,但是由于每个维度的周期都不同,整个位置编码向量的组合仍然是唯一的
- 例如,尽管单独看正弦函数sin(x)和cos(x)在每个2π的位置上的值是相同的,但是如果我们考虑函数组合sin(x)和sin(2x),那么它们在x=2π和x=4π的位置上的值是不同的
- 因此,这种编码方式可以确保即使序列的长度超过训练数据中的最大长度,每个位置的编码向量仍然是唯一的
- 因此从整体唯一性的角度,我们可以推导出个体与整体之间的绝对位置关系条件满足
- 这里实现的方式就是通过如下的公式体现周期性,因为sin 和 cos 函数自身自带周期特征
- 因此我们只需要确保每一个序列个体的高维向量中,每一个维度都具备自己的周期性即可
- pos 代表序列中的个体的index位置,如1/1000,500/1000
- i 代表序列中的个体,其PE 向量,也就是高维向量中的第i维,也就是数组中的第i个元素
- 根据欧拉公式,我们知道
- 上述公式可以转化为
- 这样,位置编码就可以看作复数平面上的一点,该点的角度是pos/10000^{2i/d_{model}},幅度是1
- 即使对于一个很大的位置pos,由于底数10000^{2i/d_{model}}的存在,所以角度值pos/10000^{2i/d_{model}}会非常小,而且随着i的增加,角度值会逐渐减小
- 每个位置的编码实际上是在复数平面上的一系列不同的点,这些点分布在以原点为中心的一系列同心圆上,每个圆对应一个不同的i
- 由于每个位置pos对应的角度是不同的,所以在每个同心圆上,每个点的位置都是不同的
- 当前公式就保证了个体在整体序列中的唯一性
- 基于数学理论,如果我们可以构建出一个关系表达式,证明函数f(x+a) = f(x) + f(a), 假设 x + a = y,那么我们可以认为 f(x) 和 f(y)之间存在相关性
- 在上面提到的公式
- 中,我们对 pos 位置的序列个体,分别使用 i 和 2i 来代表其PE高维向量中每一个维度的数学计算公式
- 如果我们可以将刚刚的结论带入,达成某种函数表达式,确保 f(pos+k) = f(pos) + f(k), 就可以证明我们需要的相对位置相关性
- 基于三角函数公式
- 借助研究学者们对数学的敏锐感知可以推导出如下
- 推导过程我就不献出我的手稿了,潦草,同学们可以自己推导一下,5分钟之内必出结果
- 上面的推导结果中,我们发现 f(pos +k) 的PE高维表示都可以通过 f(pos) 和 f(k)的PE高维表示组合而成
- 这就代表着基于上述公式,可以达成序列中不同个体的相对位置信息表达
在这篇笔记中,我们详细描述了
- 为什么 Transformer 需要 PE
- 为什么 Attention 解决不了位置问题
- PE 的设计目标
- 为什么不使用简单的方式实现PE
- Transformer 中的PE设计
- Transformer 中的PE为什么既满足绝对位置性,又满足相对位置性
关于函数式唯一性的函数图像证明及代码
import numpy as np
import matplotlib.pyplot as plt
def pos_encoding(pos, i, d_model):
return np.exp(1j * pos / np.power(10000, 2 * i / d_model))
pos = np.arange(5000) # 我们生成了0到4999的位置值
i = np.arange(512) # 这是一个示例,假设d_model=512,我们生成了0到511的维度值
pe = pos_encoding(pos[:, np.newaxis], i[np.newaxis, :], 512)
angles = np.angle(pe) # 计算角度
magnitudes = np.abs(pe) # 计算幅度
plt.figure(figsize=(12, 6))
plt.subplot(1, 2, 1)
plt.imshow(angles, aspect='auto')
plt.colorbar()
plt.title('Angles')
plt.subplot(1, 2, 2)
plt.imshow(magnitudes, aspect='auto')
plt.colorbar()
plt.title('Magnitudes')
plt.tight_layout()
plt.show()