Lorn3's blog

Transformer中的位置编码的探索与思考

Motivation

最近在完成CS336的Assignment1,从头搭建一个大模型,其中实现使用的是主流的旋转位置编码(RoPE)。虽然对照指导书完成了相关实验,但是对位置编码的选择以及原理存在一定的知识缺漏,因此写本博客以加强理解。

Note

我们在下面的数学推到中使用列向量推导

为什么Transformer需要位置编码?

设输入序列长度为L,模型的维度为d,那么我们输入的序列第i个token的表示为xid,将整个序列拼接起来我们有:

X=[x1,x2,,xL]d×L

那么Transformer Block中的MHA(由于最终结果是每一个Head通过相同的过程得到的不同结果进行concat,因此我们这里只考虑一个Head)在计算Attention Score时首先通过线性映射得到:

Q=WQXK=WKXV=WVX

其中:

WQ,WK,WVd×d

因此:

Q,K,Vd×L

接下来会计算第i个query和第j个key的相似度:

sij=qikjd

写成矩阵形式就是:

S=QKdL×L

然后进行Softmax操作(沿key维度):

A=softmax(S)L×L

最后进行加权求和:

Attention(Q,K,V)=VAd×L

PL×L为置换矩阵,表示对序列顺序的重排。我们有:

X=XP

那么我们重新计算Attention:

Q=WQX=WQXP=QPK=WKX=KPV=VPS=QKd=PQKPdA=softmax(S)=PAPAttention(Q,K,V)=VPA=VPPAP=VAP=Attention(Q,K,V)P

上式表明,在不引入任何位置信息的情况下,Attention计算具有置换同变性。但是我们需要注意到,P是一个置换矩阵,它只是索引的重定向,我们计算出来的Attention的数值没有发生任何变化。如果我们把Attention输出结果看作带有语义信息的Embedding向量,那么它的语义将不会携带位置语义的信息。

举一个具体的例子:

我们输入的文本序列是["猫","吃","鱼"],对应的输入序列为[x_1,x_2,x_3],而对应的Attention输出为[y1,y2,y3],而我们进行交换得到:["鱼","吃","猫"],对应的输入序列为[x_3,x_2,x_1],根据我们上面的证明,Attention的输出将是[y3,y2,y1],在数值上没有任何变换,只是简单的位置交换,也就是说Attention的输出没有提取出在语义结构(如主谓宾等)的信息。

位置编码为什么可以解决这个问题?

我们现在需要解决的问题是打破这种数值不变性,也就是交换了输入位置后,我们希望最终的输出为[y3',y2',y1'].我们将Attention视为一个输出为集合的函数f,那么上述不变性用数学语言表述就是:

f(,xm,,xn,)=f(,xn,,xm,)

因此,我们要做的事情就是打破这种不变性,比如在每个位置都加上一个不同的编码向量:

f^(,xm,,xn,)=f(,xm+pm,,xn+pn,)

一般来说,只要每个位置的编码向量不同,那么这种全对称性就被打破了,即可以用f~代替f来处理有序的输入。

我们写作矩阵的形式:

我们将位置编码矩阵定义为E=[e1,e2,,eL]d×L。每一个列向量ei仅与位置索引i有关。

融入后的输入矩阵为:

Xpos=X+E

此时我们重新推导Q,K,V的生成过程(以Q为例):

Q=WQ(X+E)=WQX+WQE=QX+QE

那么同理:

K=KX+KE,V=VX+VE

我们可以将{K,Q,V}E视为位置信息的“特征表达”

接下来我们计算S=QKd,我们具体来看其中一个点积项sij:

sij=1d(qx,i+qe,i)(kx,j+ke,j)=1d(qx,ikx,j+qx,ike,j+qe,ikx,j+qe,ike,j)

那么此时就多出了位置-内容,位置-位置的信息。

我们再按照同样的推理,令X=XP,那么此时:

Q=WQ(X+E)=WQ(XP+E)=QXP+QE

同理:

K=KXP+KEV=VXP+VE

那么:

S=QKd=(QXP+QE)(KXP+KE)d=1d(PQXKXP+PQXKE+QEKXP+QEKE)

显而易见,SPSP,且其中的QEKE项没有被P作用,也就意味着原本的置换对称性也就被打破了,因此我们可以用f^来代替f来处理有序的输入。

也就是说,位置编码的引入可以解决我们说的Attention置换数值不变性的问题。

怎样的位置编码是好的?

我们现在想要进一步分析位置编码的性质,从而设计更好的位置编码。我们将f^展开至二阶项(为了简化考虑,写作矩阵形式):

f^(X)=f(X+E)f(X)+f(X)·E+12EHf(X)E

那么我们来看与位置编码有关的项:

对于一阶项,只要位置编码在每个位置都是独特的就能表述绝对位置的信息。而对于二阶信息,对于一个理想的位置编码而言,应该让这个二阶项满足某种平移不变性,即对于任意位移k,位置i与位置i+k的交互模式应该是稳定的。

我们先从最简单的情况入手,假设𝐇=I为单位矩阵,那么此时EE是两个位置编码的内积,我们希望在这个简单的例子中该项表达的是相对位置信息,即存在某个函数g使得:

<pm,pn>=g(mn)

这里的pm,pn为d维向量,这里我们从最简单的d=2入手。我们称上式是一个位置编码为一个合理位置编码的条件式。

对于2维向量,我们借助复数来推导,视向量[x,y]为复数x+yi,那么我们有:

<pm,pn>=axbx+ayby=Re(pmp^n)

其中w^w的共轭复数。

为了满足上式,我们可以假设存在复数qmn,使得:

pmp^n=qmn

这样两边取实部就得到条件式。为了解这个方程,我们可以使用复数的指数形式,假设pm=rmeiϕm,p^n=rneiϕn,qmn=RmneiΦmn,那么有:

rmrnei(ϕmϕn)=RmneiΦmn

于是我们得到等式:

{rmrn=Rmnϕmϕn=Φmn

对于第一个方程,带入m = n,可以得到rm2=R0,即rm为一个常数,为了简单,我们令其为1;

对于第二个方程,显然等差数列满足上述性质(令m = 0有Φm=ϕm),设公差为θ,则通项为ϕm=Φm=mθ,由此我们得到二维情况下的位置编码的解:

pm=eimθpm=(cosmθ\newlinesinmθ)

由于内积满足线性叠加性,我们可以由二维的情况直接扩展到更高偶数维的情况:

pm=(eimθ0eimθ1eimθd/21)pm=(cosmθ0sinmθ0cosmθ1sinmθ1cosmθd/21sinmθd/21)

这样我们就求出了满足条件式的一组解,显然解不唯一。

此外,一个好的位置编码应该满足远程衰减的性质,即随着|mn|的增大,<pm,pn>有趋于0的趋势。

那么有:

<pm,pn>=Re[ei(mn)θ0+ei(mn)θ1++ei(mn)θd/21]=j=0d/21cos(kθj)k=mn

由于在LM中d通常为768,是一个较大值,因此我们可以将离散的索引j映射到连续变量t[0,1]上。设θj是某个光滑单调函数f(t)生成的,即θj=f(2j/d)。利用Euler-Maclaurin的一阶近似,我们可以将求和转化为积分:

<pm,pn>d201cos(k·f(t))dt

那么现在的问题就转化为了,寻找一个函数f(t),使得上述震荡积分在k很大时具有较好的衰减性质(足够快)。

根据黎曼-勒贝格引理,只要频率分布函数f(t)满足光滑且严格单调的条件,当kinf时,被积函数的高频震荡将导致正负面积互相抵消,使得积分值必然趋近于0.

在Transformer2017的论文中的Sinusoidal位置编码选择的是θt=10000t.

由此,我们便推导出了Sinusoidal位置编码的形式:

{pk,2i=sin(k/100002i/d)pk,2i+1=cos(k/100002i/d)

我们只能说明它的合理性,但是无法说明它的最优性,因为它并不一定最优[Lol]

事实上,一个可行的方案是将位置编码中的θi设为可学习的参数,其初始值为θi=100002i/d.

在上述推导中,都是基于H = I这个简单情况,对于一般的H,使用上述Sinusoidal位置编码,还能具备我们理想的性质吗?

事实上,有研究表明[3] 在网络规模足够大的情况下,其Hessian矩阵将呈现出块对角的形式,因此我们考虑H是一个对角阵的情况,此时:

pmHpn=i=1d/2H2i,2icosmθicosnθi+H2i+1,2i+1sinmθisinnθi

由和差化积有:

i=1d/212(H2i,2i+H2i+1,2i+1)cos(mn)θi+12(H2i,2iH2i+1,2i+1)cos(m+n)θi

可以看到其中是包含了相对位置项(m-n)的,只是会出现m+n项。

因此我们可以认为Sinusoidal位置编码是一个有效的位置编码。

为什么RoPE[4]比Sinusoidal位置编码更好?

在先前分析为什么需要位置编码的过程中,我们知道位置编码的作用是在计算Attention Score时,让QK点积中包含文本的结构语义信息,具体而言是下式中的后三项包含位置信息:

S=QKd=(QXP+QE)(KXP+KE)d=1d(PQXKXP+PQXKE+QEKXP+QEKE)

我们可以发现,除了最后一项是相对位置信息外,另外两项则是绝对位置信息与文本语义信息的耦合,这实际上是一种冗余,因为绝对位置信息的作用相较于相对位置信息而言一方面具有误导性,例如在短文本中结尾出现在绝对位置10中,而10在长文本中可能仅仅是文本开头;另一方面模型还会额外去学习已经得到的相对位置信息造成资源的浪费.

为了解决这个问题,工业界进行了许多尝试[5],但是这些尝试普遍会带来大量的额外计算存储开销。由于绝对位置编码具有实现简单,计算速度快的特点,并且在Sinusoidal位置编码中我们也能看到通过绝对位置编码在一定程度上是可以得到相对位置信息的,如果可以通过绝对位置编码的方式实现相对位置编码,那么就是“集各家之所长”,“鱼和熊掌兼得”了。

为了实现这个目标,我们假设通过下述运算来给q,k添加绝对位置信息:

q^m=f(q,m)k^n=f(k,n)

而我们希望得到如下恒等关系:

<f(q,m),f(k,n)>=g(q,k,mn)

为了求解方便,我们令:

f(q,0)=qf(k,0)=k

求解思路与之前推导Sinusoidal类似,我们先考虑二维的情况,然后借助复数来求解。

Re[f(q,m)f^(k,n)]=g(q,k,mn)

设:

f(q,m)=Rf(q,m)eiθf(q,m)f^(k,n)=Rf(k,n)eiθf(k,n)g(q,k,mn)=Rg(q,k,mn)eiθg(q,k,mn)

那么带入方程求解得到:

{Rf(q,m)Rf(k,n)=Rg(q,k,mn)θf(q,m)θf(k,n)=θg(q,k,mn)

对于第一个方程,令m = n有:

Rf(q,m)Rf(k,n)=Rg(q,k,0)=Rf(q,0)Rf(k,0)=||q||||k||

那么我们可以直接令Rf(q,m)=||q||,即它不依赖于m。

对于第二个方程,同样地令m=n有:

θf(q,m)θf(k,n)=θg(q,k,0)=θf(q,0)θf(k,0)=θqθk

这里的θq,θk是q,k本身的辐角。

由此我们得到:

θf(q,m)θq=θf(k,n)θk

因此θf(q,m)θq应该是一个只与m有关的而与q无关的函数,因为我们希望恒等式对任意的q,k都成立,所以右侧必须与q,k无关,因此记为常数,记为θ,该式记为ϕ(m),即

θf(q,m)=θq+ϕ(m)

令n = m - 1 有:

ϕ(m)ϕ(m1)=θg(q,k,1)+θkθq

{ϕ(m)}为等差数列,公差为θ,得到ϕ(m)=mθ

由此,我们得到了二维情况下用复数表示的RoPE:

f(q,m)=Rf(q,m)eiθf(q,m)=||q||ei(θq+mθ)=qeimθ

根据复数乘法的集合意义,该变换实际上对应着向量的旋转,我们可以将其写为矩阵形式:

f(q,m)=(cosmθsinmθsinmθcosmθ)(q0q1)

由内积的线性叠加性,我们可以得到任意偶数维度的RoPE:

(cosmθ0sinmθ00000sinmθ0cosmθ0000000cosmθ1sinmθ1000000cosmθd/21sinmθd/210000sinmθd/21cosmθd/21)(q0q1q2q3qd2qd1)

我们称左侧的矩阵为旋转矩阵记为Rm.也就是说给位置为m的向量q乘上矩阵Rm,位置为n的向量k乘上矩阵Rn,用变换后的Q,K序列做Attention,那么Attention就自动包含相对位置信息了,因为恒等式成立:

(Rmq)(Rnk)=qRmRnk=qRnmk

值得指出的是,Rm是一个正交矩阵,它不会改变向量的模长,因此通常来说它不会改变原模型的稳定性。

此外,在具体实现时,我们并不会拿这个大矩阵去乘以向量,而是采用逐位相乘的方式:

(q0q1q2q3qd2qd1)(cosmθ0cosmθ0cosmθ1cosmθ1cosmθd/21cosmθd/21)+(q1q0q3q2qd1qd2)(sinmθ0sinmθ0sinmθ1sinmθ1sinmθd/21sinmθd/21)

对于θ的选择,作者选择了与Sinusoidal位置编码一样的θi=100002i/d,从而带来远程衰减性。

至此,我们完成了对RoPE的推导,并通过推导成功说明了为何RoPE相较于Sinusoidal位置编码更好,因为它通过绝对位置的注入方式,实现了相对位置的注入,而不带来其它冗余。

最后,下面是我在CS336中实现的一个RoPE:

from einops import rearrange, einsum
import torch
import torch.nn as nn

class RotaryPositionalEmbedding(nn.Module):
    def __init__(self,theta,d_k,max_seq_len,device):
        """
        d_k: int, 维度大小,必须为偶数
        theta: float, RoPE中的\Theta值
        max_seq_len: int, 最大序列长度
        device: torch.device, 设备
        """
        super().__init__()
        assert d_k % 2 == 0, "d_k must be even"
        self.theta = theta
        self.d_k = d_k
        self.max_seq_len = max_seq_len
        self.device = device

        #一共有d / 2个频率
        half_dk = d_k // 2
        k = torch.arange(0,half_dk,device=device).float()
        inv_freq = 1.0 / (self.theta ** (2.0 * k / d_k))

        positions = torch.arange(0,max_seq_len,device = device).float()

        angles = einsum(positions,inv_freq,"max_seq_len,half_dk->max_seq_len half_dk")
        cos = torch.cos(angles)
        sin = torch.sin(angles)

        self.register_buffer("cos",cos,persistent = False)
        self.register_buffer("sin",sin,persistent = False)

    def forward(self,x,token_positions):
        """
        inputs:
            x: ...,seq_len,d_k
            token_positions:...,seq_len
        returns:
            x_rotated: ...,seq_len,d_k
        """
        cos = self.cos[token_positions]  # ...,seq_len,half_dk
        sin = self.sin[token_positions]  # ...,seq_len,half_dk

        x_even = x[...,0::2]
        x_odd = x[...,1::2]

        x_rot_even = x_even * cos - x_odd * sin
        x_rot_odd = x_even * sin + x_odd * cos

        out = torch.empty_like(x)
        out[...,0::2] = x_rot_even
        out[...,1::2] = x_rot_odd
        return out

参考:

  1. Transformer升级之路:1、Sinusoidal位置编码追根溯源 - 科学空间|Scientific Spaces
  2. Transformer升级之路:2、博采众长的旋转式位置编码 - 科学空间|Scientific Spaces
  3. [2505.02809] Towards Quantifying the Hessian Structure of Neural Networks
  4. RoFormer: Enhanced Transformer with Rotary Position Embedding | Cool Papers - Immersive Paper Discovery
  5. 让研究人员绞尽脑汁的Transformer位置编码 - 科学空间|Scientific Spaces