RoPE原理

不讲从0-1的推导原理

背景是Transformerpositional embedding直接添加在word_embedding

[12]⏟word−embedding+[0.30.6]⏟positional−embedding=[1.32.6] \underbrace{\begin{bmatrix} 1\\ 2\end{bmatrix}}_{word-embedding}+\underbrace{\begin{bmatrix} 0.3\\ 0.6\end{bmatrix}}_{positional-embedding}=\begin{bmatrix} 1.3\\ 2.6\end{bmatrix} wordembedding [12]+positionalembedding [0.30.6]=[1.32.6]
发现对word_embedding的模长和角度都发生了变化,虽然有一定程度的位置区分度,但是也很大程度上引入了一部分噪声

Q=Wq(Xm+Pm)K=Wk(Xn+Pn)QKT≈(Xm+Pm)(Xn+Pn)T=XmXnT⏟词向量间+XmPnT+PmXnT⏟噪声+PmPnT⏟位置编码间 \begin{aligned} Q&= W_q(X_m+P_m)\\ K&= W_k(X_n+P_n)\\ QK^T&\approx (X_m+P_m)(X_n+P_n)^T\\ &=\underbrace{X_mX^T_n}_{词向量间}+\underbrace{X_mP^T_n+P_mX^T_n}_{噪声}+\underbrace{P_mP^T_n}_{位置编码间} \end{aligned} QKQKT=Wq(Xm+Pm)=Wk(Xn+Pn)(Xm+Pm)(Xn+Pn)T=词向量间 XmXnT+噪声 XmPnT+PmXnT+位置编码间 PmPnT

因此引入RoPE,一种将位置信息以向量旋转的方式,不改变模长的条件下进行引入的方式

首先定义q@[bs,seq,n_head,head_dim]k/v@[bs,seq,n_kv_head,head_dim]

基于

(a⏟real+ib⏟imaginary)(cos⁡θ+isin⁡θ)=(acos⁡θ−bsin⁡θ)⏟real−after−rotateθ+i(asin⁡θ+bcos⁡θ)⏟imaginary−after−rotateθ (\underbrace{a}_{real}+i\underbrace{b}_{imaginary})(\cos\theta+i\sin\theta)=\underbrace{(a\cos\theta-b\sin\theta)}_{real-after-rotate \theta}+i\underbrace{(a\sin\theta+b\cos\theta)}_{imaginary-after-rotate \theta} (real a+iimaginary b)(cosθ+isinθ)=realafterrotateθ (acosθbsinθ)+iimaginaryafterrotateθ (asinθ+bcosθ)

这个公式,实现如下的向量旋转

[ab]→[acos⁡θ−bsin⁡θasin⁡θ+bcos⁡θ] \begin{bmatrix} a\\ b\end{bmatrix} \rightarrow\begin{bmatrix} a\cos\theta-b\sin\theta\\ a\sin\theta+b\cos\theta\end{bmatrix} [ab][acosθbsinθasinθ+bcosθ]

所以也可以通过q−>qr+jqi−>qr′+jqi′=>q′q->q_r +j q_i ->q'_r +j q'_i=>q'q>qr+jqi>qr+jqi=>q这个方式实现对于query向量的旋转

因此有q_r@[bs,seq,n_head,head_dim//2]以及q_i@[bs,seq,n_head,head_dim//2]的拆分

[qr(1,1)⋯qr(1,d′)⋯qr(m,k)⋯qr(seq,1)⋯qr(seq,d′)]seq×d′⏟qr∗f([1⋅10000−2⋅1d⋯1⋅10000−2⋅d′d⋯cos⁡(mθk)⋯seq⋅10000−2⋅1d⋯seq⋅10000−2⋅d′d]seq×d′⏟freq) \underbrace{\begin{bmatrix} q^{(1,1)}_r &\cdots &q^{(1,d')}_r\\ \cdots &q^{(m,k)}_r &\cdots \\ q^{(seq,1)}_r &\cdots &q^{(seq,d')}_r\end{bmatrix}_{seq\times d'}}_{q_r} * f(\underbrace{\begin{bmatrix} 1\cdot10000^{-\frac{2\cdot1}{d}} &\cdots &1\cdot10000^{-\frac{2\cdot d'}{d}}\\ \cdots &\cos(m\theta_k) &\cdots \\ seq\cdot10000^{-\frac{2\cdot1}{d}} &\cdots &seq\cdot10000^{-\frac{2\cdot d'}{d}}\end{bmatrix}_{seq\times d'}}_{freq}) qr qr(1,1)qr(seq,1)qr(m,k)qr(1,d)qr(seq,d) seq×df(freq 110000d21seq10000d21cos(mθk)110000d2dseq10000d2d seq×d)
其中d′=d2d'=\frac{d}{2}d=2d

相当于实现了

qr′=qr∗cos⁡(freq)−qi∗sin⁡(freq)qi′=qr∗sin⁡(freq)+qi∗cos⁡(freq) \begin{aligned} q'_r&=q_r*\cos(freq)-q_i*\sin(freq)\\ q'_i&=q_r*\sin(freq)+q_i*\cos(freq) \end{aligned} qrqi=qrcos(freq)qisin(freq)=qrsin(freq)+qicos(freq)

苏神在博客中写道

[cos⁡mθ0−sin⁡mθ000⋯00sin⁡mθ0cos⁡mθ000⋯0000cos⁡mθ1−sin⁡mθ1⋯0000sin⁡mθ1cos⁡mθ1⋯00⋯⋯⋯⋯⋯⋯⋯0000⋯cos⁡mθd/2−1−sin⁡mθd/2−10000⋯sin⁡mθd/2−1cos⁡mθd/2−1]⏟Rm[q0q1q2q3⋯qd−2qd−1] \underbrace{\begin{bmatrix} \cos m\theta_0 &-\sin m\theta_0 &0 &0 &\cdots &0 &0\\ \sin m\theta_0 &\cos m\theta_0 &0 &0 &\cdots &0 &0\\ 0 & 0 &\cos m\theta_1 &-\sin m\theta_1 &\cdots &0 &0\\ 0 &0 &\sin m\theta_1 &\cos m\theta_1 &\cdots &0 &0\\ \cdots &\cdots &\cdots &\cdots &\cdots &\cdots &\cdots\\ 0 &0 &0 &0&\cdots &\cos m\theta_{d/2-1} &-\sin m\theta_{d/2-1} \\0 &0 &0 &0&\cdots &\sin m\theta_{d/2-1} &\cos m\theta_{d/2-1} \end{bmatrix}}_{R_m} \begin{bmatrix} q_0\\ q_1\\ q_2\\ q_3\\ \cdots \\ q_{d-2} \\q_{d-1}\end{bmatrix} Rm cosmθ0sinmθ00000sinmθ0cosmθ0000000cosmθ1sinmθ10000sinmθ1cosmθ1000000cosmθd/21sinmθd/210000sinmθd/21cosmθd/21 q0q1q2q3qd2qd1
并且定义

[q0′q1′]=[cos⁡mθ0−sin⁡mθ0sin⁡mθ0cos⁡mθ0][q0q1][q2′q3′]=[cos⁡mθ1−sin⁡mθ1sin⁡mθ1cos⁡mθ1][q2q3]⋯ \begin{aligned} \begin{bmatrix} q'_0\\ q'_1 \end{bmatrix}&= \begin{bmatrix} \cos m\theta_0 &-\sin m\theta_0 \\ \sin m\theta_0 &\cos m\theta_0 \end{bmatrix} \begin{bmatrix} q_0\\ q_1 \end{bmatrix}\\ \begin{bmatrix} q'_2\\ q'_3 \end{bmatrix}&= \begin{bmatrix} \cos m\theta_1 &-\sin m\theta_1 \\ \sin m\theta_1 &\cos m\theta_1 \end{bmatrix} \begin{bmatrix} q_2\\ q_3 \end{bmatrix}\\ &\cdots \end{aligned} [q0q1][q2q3]=[cosmθ0sinmθ0sinmθ0cosmθ0][q0q1]=[cosmθ1sinmθ1sinmθ1cosmθ1][q2q3]
去实现q=[q0 q1 ... q_{d-1}] -> q'=[q'0 q'1 ... q'{d-1}]旋转

[!important]

这里容易误解的点在于,苏神是对f(q,m)=[cos⁡mθ−sin⁡mθsin⁡mθcos⁡mθ][q0q1]f(q,m)=\begin{bmatrix}\cos m\theta & -\sin m\theta \\ \sin m\theta & \cos m\theta\end{bmatrix}\begin{bmatrix}q_0 \\ q_1\end{bmatrix}f(q,m)=[cosmθsinmθsinmθcosmθ][q0q1]的一个特例一定要理解的是m表示的是当前的token在整个sequence中的位置。所以这里相当于

q=[q(1,1)⋯q(1,d)⋯q(i,j)⋯q(seq,1)⋯q(seq,d)]seq×d q=\begin{bmatrix} q^{(1,1)} &\cdots &q^{(1,d)}\\ \cdots &q^{(i,j)} &\cdots \\ q^{(seq,1)} &\cdots &q^{(seq,d)} \end{bmatrix}_{seq\times d} q= q(1,1)q(seq,1)q(i,j)q(1,d)q(seq,d) seq×d

在第m行,把一行数据拿出来

[q(m,1)q(m,2)q(m,3)⋯q(m,d)]1×d \begin{bmatrix} q^{(m,1)} &q^{(m,2)} &q^{(m,3)} &\cdots &q^{(m,d)} \end{bmatrix}_{1\times d} [q(m,1)q(m,2)q(m,3)q(m,d)]1×d

并把它称作为[q0,q1,q2,⋯ ,qd−1][q_0, q_1, q_2, \cdots, q_{d-1}][q0,q1,q2,,qd1]

因此不要忘记要对除m行以外的所有q(i,j)q^{(i,j)}q(i,j)都进行上述的旋转操作!

因此,推导到这里,实际上,这和我上文就是一致的了:

我们从元素级别入手:苏神在博客中的q0∈Rbs×1×n−head×1q_0 \in \mathbb{R}^{bs\times 1\times n-head\times 1}q0Rbs×1×nhead×1,切记这里的q0q_0q0实则是q(m,0)q^{(m,0)}q(m,0),和它交互的是cos⁡mθ0\cos m\theta_0cosmθ0sin⁡mθ0\sin m\theta_0sinmθ0;而我们的推导当中qr(m,k)∈Rbs×1×n−head×1q^{(m,k)}_r\in \mathbb{R}^{bs \times 1 \times n-head \times 1}qr(m,k)Rbs×1×nhead×1,与它交互的是cos⁡mθk\cos m\theta_kcosmθksin⁡mθk\sin m\theta_ksinmθk。从元素级别可以看到每一个元素的交互都是一样的。

[!important]

事实上,拓展到高维后,有结论如下:

对于任意维度ddd,任意第kkkk=0,1,...,d/2−1k=0,1,...,d/2-1k=0,1,...,d/21

q2k=qr[k]q_{2k} = q_r[k]q2k=qr[k]第 k 组的第一个元素 = 第 k 个复数的实部

q2k+1=qi[k]q_{2k+1} = q_i[k]q2k+1=qi[k]第 k 组的第二个元素 = 第 k 个复数的虚部

我们可以把整体进行一个缩小,由于各方法对q@[bs,seq,n_head,head_dim]bs,n_head不改动,后续高维推导中省略这2个维度,所以这里q∈Rseq×d\boldsymbol{q}\in \mathbb{R}^{seq\times d}qRseq×d简化为如下所示

q=[q(0,0)q(0,1)⋯q(0,d−1)q(1,0)q(1,1)⋯q(1,d−1)⋯⋯⋯⋯q(seq−1,0)q(seq−1,1)⋯q(seq−1,d−1)] \boldsymbol{q}=\begin{bmatrix} q^{(0,0)} & q^{(0,1)} &\cdots &q^{(0,d-1)}\\ q^{(1,0)} & q^{(1,1)} &\cdots &q^{(1,d-1)}\\ \cdots &\cdots &\cdots &\cdots \\ q^{(seq-1,0)} & q^{(seq-1,1)} &\cdots &q^{(seq-1,d-1)}\\ \end{bmatrix} q= q(0,0)q(1,0)q(seq1,0)q(0,1)q(1,1)q(seq1,1)q(0,d1)q(1,d1)q(seq1,d1)

根据上述结论可知:

qr=[q(0,0)q(0,2)⋯q(0,d−2)q(1,0)q(1,2)⋯q(1,d−2)⋯⋯⋯⋯q(seq−1,0)q(seq−1,2)⋯q(seq−1,d−2)]qi=[q(0,1)q(0,3)⋯q(0,d−1)q(1,1)q(1,3)⋯q(1,d−1)⋯⋯⋯⋯q(seq−1,1)q(seq−1,3)⋯q(seq−1,d−1)] \begin{aligned} \boldsymbol{q_r}&=\begin{bmatrix} q^{(0,0)} & q^{(0,2)} &\cdots &q^{(0,d-2)}\\ q^{(1,0)} & q^{(1,2)} &\cdots &q^{(1,d-2)}\\ \cdots &\cdots &\cdots &\cdots \\ q^{(seq-1,0)} & q^{(seq-1,2)} &\cdots &q^{(seq-1,d-2)}\\ \end{bmatrix}\\ \boldsymbol{q_i}&=\begin{bmatrix} q^{(0,1)} & q^{(0,3)} &\cdots &q^{(0,d-1)}\\ q^{(1,1)} & q^{(1,3)} &\cdots &q^{(1,d-1)}\\ \cdots &\cdots &\cdots &\cdots \\ q^{(seq-1,1)} & q^{(seq-1,3)} &\cdots &q^{(seq-1,d-1)}\\ \end{bmatrix}\\ \end{aligned} qrqi= q(0,0)q(1,0)q(seq1,0)q(0,2)q(1,2)q(seq1,2)q(0,d2)q(1,d2)q(seq1,d2) = q(0,1)q(1,1)q(seq1,1)q(0,3)q(1,3)q(seq1,3)q(0,d1)q(1,d1)q(seq1,d1)

写到这里,也更加清晰,苏神的推导中取q\boldsymbol{q}qmmm行进行RmR_mRm矩阵旋转:

q′(m,2i)=q(m,2i)⋅cos⁡mθi−q(m,2i+1)⋅sin⁡mθiq′(m,2i+1)=q(m,2i)⋅sin⁡mθi+q(m,2i+1)⋅cos⁡mθi \begin{aligned} q'^{(m,2i)}&=q^{(m,2i)}\cdot \cos m\theta_i - q^{(m,2i+1)}\cdot \sin m\theta_i \\ q'^{(m,2i+1)}&=q^{(m,2i)}\cdot \sin m\theta_i + q^{(m,2i+1)}\cdot \cos m\theta_i \end{aligned} q(m,2i)q(m,2i+1)=q(m,2i)cosmθiq(m,2i+1)sinmθi=q(m,2i)sinmθi+q(m,2i+1)cosmθi

而工程代码中总结规律发现上述方法效率太低,既然:

  1. q(m,2i)q^{(m,2i)}q(m,2i)q(m,2i+1)q^{(m,2i+1)}q(m,2i+1)总是有固定的cos⁡mθi\cos m\theta_icosmθi/sin⁡mθi\sin m\theta_isinmθi−sin⁡mθi-\sin m\theta_isinmθi/cos⁡mθi\cos m\theta_icosmθi要相乘;
  2. q(m,2i)q^{(m,2i)}q(m,2i)q(m,2i+1)q^{(m,2i+1)}q(m,2i+1)恰好对应qr\boldsymbol{q_r}qrqi\boldsymbol{q_i}qi中的一行;

不如直接整合成一个矩阵,实现

[q′(0,0)q′(0,2)⋯q′(0,d−2)q′(1,0)q′(1,2)⋯q′(1,d−2)⋯⋯⋯⋯q′(seq−1,0)q′(seq−1,2)⋯q′(seq−1,d−2)]⏟qr′=[q(0,0)q(0,2)⋯q(0,d−2)q(1,0)q(1,2)⋯q(1,d−2)⋯⋯⋯⋯q(seq−1,0)q(seq−1,2)⋯q(seq−1,d−2)]⏟qr∗cos⁡([0θ00θ1⋯0θd/2−11θ01θ1⋯1θd/2−1⋯⋯⋯⋯(seq−1)θ0(seq−1)θ1⋯(seq−1)θd/2−1])⏟cos⁡mθ−[q(0,1)q(0,3)⋯q(0,d−1)q(1,1)q(1,3)⋯q(1,d−1)⋯⋯⋯⋯q(seq−1,1)q(seq−1,3)⋯q(seq−1,d−1)]⏟qi∗sin⁡([0θ00θ1⋯0θd/2−11θ01θ1⋯1θd/2−1⋯⋯⋯⋯(seq−1)θ0(seq−1)θ1⋯(seq−1)θd/2−1])⏟sin⁡mθ[q′(0,1)q′(0,3)⋯q′(0,d−1)q′(1,1)q′(1,3)⋯q′(1,d−1)⋯⋯⋯⋯q′(seq−1,1)q′(seq−1,3)⋯q′(seq−1,d−1)]⏟qi′=[q(0,0)q(0,2)⋯q(0,d−2)q(1,0)q(1,2)⋯q(1,d−2)⋯⋯⋯⋯q(seq−1,0)q(seq−1,2)⋯q(seq−1,d−2)]⏟qr∗sin⁡([0θ00θ1⋯0θd/2−11θ01θ1⋯1θd/2−1⋯⋯⋯⋯(seq−1)θ0(seq−1)θ1⋯(seq−1)θd/2−1])⏟sin⁡mθ+[q(0,1)q(0,3)⋯q(0,d−1)q(1,1)q(1,3)⋯q(1,d−1)⋯⋯⋯⋯q(seq−1,1)q(seq−1,3)⋯q(seq−1,d−1)]⏟qi∗cos⁡([0θ00θ1⋯0θd/2−11θ01θ1⋯1θd/2−1⋯⋯⋯⋯(seq−1)θ0(seq−1)θ1⋯(seq−1)θd/2−1])⏟cos⁡mθ \begin{aligned} \underbrace{\begin{bmatrix} q'^{(0,0)} & q'^{(0,2)} &\cdots &q'^{(0,d-2)}\\ q'^{(1,0)} & q'^{(1,2)} &\cdots &q'^{(1,d-2)}\\ \cdots &\cdots &\cdots &\cdots \\ q'^{(seq-1,0)} & q'^{(seq-1,2)} &\cdots &q'^{(seq-1,d-2)}\\ \end{bmatrix}}_{\boldsymbol{q'_r}}&= \underbrace{\begin{bmatrix} q^{(0,0)} & q^{(0,2)} &\cdots &q^{(0,d-2)}\\ q^{(1,0)} & q^{(1,2)} &\cdots &q^{(1,d-2)}\\ \cdots &\cdots &\cdots &\cdots \\ q^{(seq-1,0)} & q^{(seq-1,2)} &\cdots &q^{(seq-1,d-2)}\\ \end{bmatrix}}_{\boldsymbol{q_r}}* \underbrace{\cos(\begin{bmatrix} 0\theta_0 &0\theta_1 &\cdots &0\theta_{d/2-1}\\ 1\theta_0 &1\theta_1 &\cdots &1\theta_{d/2-1}\\ \cdots &\cdots &\cdots &\cdots \\ (seq-1)\theta_0 &(seq-1)\theta_1 &\cdots &(seq-1)\theta_{d/2-1}\\ \end{bmatrix})}_{\boldsymbol{\boldsymbol{\cos m\theta}}}- \underbrace{\begin{bmatrix} q^{(0,1)} & q^{(0,3)} &\cdots &q^{(0,d-1)}\\ q^{(1,1)} & q^{(1,3)} &\cdots &q^{(1,d-1)}\\ \cdots &\cdots &\cdots &\cdots \\ q^{(seq-1,1)} & q^{(seq-1,3)} &\cdots &q^{(seq-1,d-1)}\\ \end{bmatrix}}_{\boldsymbol{q_i}}* \underbrace{\sin(\begin{bmatrix} 0\theta_0 &0\theta_1 &\cdots &0\theta_{d/2-1}\\ 1\theta_0 &1\theta_1 &\cdots &1\theta_{d/2-1}\\ \cdots &\cdots &\cdots &\cdots \\ (seq-1)\theta_0 &(seq-1)\theta_1 &\cdots &(seq-1)\theta_{d/2-1}\\ \end{bmatrix})}_{\boldsymbol{\boldsymbol{\sin m\theta}}} \\ \underbrace{\begin{bmatrix} q'^{(0,1)} & q'^{(0,3)} &\cdots &q'^{(0,d-1)}\\ q'^{(1,1)} & q'^{(1,3)} &\cdots &q'^{(1,d-1)}\\ \cdots &\cdots &\cdots &\cdots \\ q'^{(seq-1,1)} & q'^{(seq-1,3)} &\cdots &q'^{(seq-1,d-1)}\\ \end{bmatrix}}_{\boldsymbol{q'_i}}&= \underbrace{\begin{bmatrix} q^{(0,0)} & q^{(0,2)} &\cdots &q^{(0,d-2)}\\ q^{(1,0)} & q^{(1,2)} &\cdots &q^{(1,d-2)}\\ \cdots &\cdots &\cdots &\cdots \\ q^{(seq-1,0)} & q^{(seq-1,2)} &\cdots &q^{(seq-1,d-2)}\\ \end{bmatrix}}_{\boldsymbol{q_r}}* \underbrace{\sin(\begin{bmatrix} 0\theta_0 &0\theta_1 &\cdots &0\theta_{d/2-1}\\ 1\theta_0 &1\theta_1 &\cdots &1\theta_{d/2-1}\\ \cdots &\cdots &\cdots &\cdots \\ (seq-1)\theta_0 &(seq-1)\theta_1 &\cdots &(seq-1)\theta_{d/2-1}\\ \end{bmatrix})}_{\boldsymbol{\boldsymbol{\sin m\theta}}}+ \underbrace{\begin{bmatrix} q^{(0,1)} & q^{(0,3)} &\cdots &q^{(0,d-1)}\\ q^{(1,1)} & q^{(1,3)} &\cdots &q^{(1,d-1)}\\ \cdots &\cdots &\cdots &\cdots \\ q^{(seq-1,1)} & q^{(seq-1,3)} &\cdots &q^{(seq-1,d-1)}\\ \end{bmatrix}}_{\boldsymbol{q_i}}* \underbrace{\cos(\begin{bmatrix} 0\theta_0 &0\theta_1 &\cdots &0\theta_{d/2-1}\\ 1\theta_0 &1\theta_1 &\cdots &1\theta_{d/2-1}\\ \cdots &\cdots &\cdots &\cdots \\ (seq-1)\theta_0 &(seq-1)\theta_1 &\cdots &(seq-1)\theta_{d/2-1}\\ \end{bmatrix})}_{\boldsymbol{\boldsymbol{\cos m\theta}}} \end{aligned} qr q(0,0)q(1,0)q(seq1,0)q(0,2)q(1,2)q(seq1,2)q(0,d2)q(1,d2)q(seq1,d2) qi q(0,1)q(1,1)q(seq1,1)q(0,3)q(1,3)q(seq1,3)q(0,d1)q(1,d1)q(seq1,d1) =qr q(0,0)q(1,0)q(seq1,0)q(0,2)q(1,2)q(seq1,2)q(0,d2)q(1,d2)q(seq1,d2) cos cos( 0θ01θ0(seq1)θ00θ11θ1(seq1)θ10θd/211θd/21(seq1)θd/21 )qi q(0,1)q(1,1)q(seq1,1)q(0,3)q(1,3)q(seq1,3)q(0,d1)q(1,d1)q(seq1,d1) sin sin( 0θ01θ0(seq1)θ00θ11θ1(seq1)θ10θd/211θd/21(seq1)θd/21 )=qr q(0,0)q(1,0)q(seq1,0)q(0,2)q(1,2)q(seq1,2)q(0,d2)q(1,d2)q(seq1,d2) sin sin( 0θ01θ0(seq1)θ00θ11θ1(seq1)θ10θd/211θd/21(seq1)θd/21 )+qi q(0,1)q(1,1)q(seq1,1)q(0,3)q(1,3)q(seq1,3)q(0,d1)q(1,d1)q(seq1,d1) cos cos( 0θ01θ0(seq1)θ00θ11θ1(seq1)θ10θd/211θd/21(seq1)θd/21 )

具体复杂度对比这里就不做了,工程代码实现了从O(n2)→O(n)O(n^2) \rightarrow O(n)O(n2)O(n)的简化

此外,由于RoPE

Q=Wq(RmXm)K=Wk(RnXn)QKT≈(RmXm)(RnXn)T=RmXmXnTRnT \begin{aligned} Q&= W_q(R_mX_m)\\ K&= W_k(R_nX_n)\\ QK^T&\approx (R_mX_m)(R_nX_n)^T\\ &=R_mX_mX^T_nR^T_n \end{aligned} QKQKT=Wq(RmXm)=Wk(RnXn)(RmXm)(RnXn)T=RmXmXnTRnT
没有引入类似positional embedding中的噪声项,因此也更加稳定。

最后由衷说一声,苏神牛逼

代码实现

def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:
    # 获取输入张量的形状:批量大小、序列长度、键/值对头的数量、每个头的维度大小
    bs, slen, n_kv_heads, head_dim = x.shape
    
    # 如果重复次数为1,则不需要重复,直接返回原始张量
    if n_rep == 1:
        return x
    
    # 对张量进行扩展和重塑操作以重复键值对
    return (
        x[:, :, :, None, :]  # 在第四个维度(头的维度前)添加一个新的维度
        .expand(bs, slen, n_kv_heads, n_rep, head_dim)  # 将新添加的维度扩展到n_rep大小,实现重复的效果
        .reshape(bs, slen, n_kv_heads * n_rep, head_dim)  # 重新塑形,合并键/值对头的数量和重复次数的维度
    )
    
# 注意:此处的dim应为 dim//n_head,因为我们是对每个head进行旋转嵌入
def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0):
    # torch.arange(0, dim, 2)[: (dim // 2)].float()生成了一个从0开始,步长为2的序列,长度为dim的一半
    # 然后每个元素除以dim,再取theta的倒数,得到频率
    freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
    # 生成一个从0到end的序列,长度为end
    t = torch.arange(end, device=freqs.device)
    # 计算外积,得到一个二维矩阵,每一行是t的元素乘以freqs的元素
    freqs = torch.outer(t, freqs).float()
    # 计算频率的余弦值,得到实部
    freqs_cos = torch.cos(freqs)
    # 计算频率的正弦值,得到虚部
    freqs_sin = torch.sin(freqs)
    return freqs_cos, freqs_sin
    
def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor):
    # 获取x的维度数
    ndim = x.ndim
    
    # 断言,确保1在x的维度范围内
    assert 0 <= 1 < ndim
    
    # 断言,确保freqs_cis的形状与x的第二维和最后一维相同
    assert freqs_cis.shape == (x.shape[1], x.shape[-1])
    
    # 构造一个新的形状,除了第二维和最后一维,其他维度都为1,这样做是为了能够将freqs_cis与x进行广播操作
    shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
    
    # 将freqs_cis调整为新的形状,并返回
    return freqs_cis.view(shape)
    
def apply_rotary_emb(
    xq: torch.Tensor,
    xk: torch.Tensor,
    freqs_cos: torch.Tensor,
    freqs_sin: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:

    # 将查询和键张量转换为浮点数,并重塑形状以分离实部和虚部
    xq_r, xq_i = xq.float().reshape(xq.shape[:-1] + (-1, 2)).unbind(-1)
    xk_r, xk_i = xk.float().reshape(xk.shape[:-1] + (-1, 2)).unbind(-1)

    # 重新塑形频率张量以进行广播
    freqs_cos = reshape_for_broadcast(freqs_cos, xq_r)
    freqs_sin = reshape_for_broadcast(freqs_sin, xq_r)

    # 应用旋转,分别计算旋转后的实部和虚部
    xq_out_r = xq_r * freqs_cos - xq_i * freqs_sin
    xq_out_i = xq_r * freqs_sin + xq_i * freqs_cos
    xk_out_r = xk_r * freqs_cos - xk_i * freqs_sin
    xk_out_i = xk_r * freqs_sin + xk_i * freqs_cos

    # 将最后两个维度合并,并还原为原始张量的形状
    xq_out = torch.stack([xq_out_r, xq_out_i], dim=-1).flatten(3)
    xk_out = torch.stack([xk_out_r, xk_out_i], dim=-1).flatten(3)

    return xq_out.type_as(xq), xk_out.type_as(xk)

class Attention(nn.Module):
    def __init__(self, args: ModelConfig):
        super().__init__()
        # 根据是否指定n_kv_heads,确定用于键(key)和值(value)的头的数量。
        self.n_kv_heads = args.n_heads if args.n_kv_heads is None else args.n_kv_heads
        # 确保总头数可以被键值头数整除。
        assert args.n_heads % self.n_kv_heads == 0

        # 模型并行处理大小,默认为1。
        model_parallel_size = 1
        # 本地计算头数,等于总头数除以模型并行处理大小。
        self.n_local_heads = args.n_heads // model_parallel_size
        # 本地键值头数,等于键值头数除以模型并行处理大小。
        self.n_local_kv_heads = self.n_kv_heads // model_parallel_size
        # 重复次数,用于扩展键和值的尺寸。
        self.n_rep = self.n_local_heads // self.n_local_kv_heads
        # 每个头的维度,等于模型维度除以头的总数。
        self.head_dim = args.dim // args.n_heads

        # 定义权重矩阵。
        self.wq = nn.Linear(args.dim, args.n_heads * self.head_dim, bias=False)
        self.wk = nn.Linear(args.dim, self.n_kv_heads * self.head_dim, bias=False)
        self.wv = nn.Linear(args.dim, self.n_kv_heads * self.head_dim, bias=False)
        # 输出权重矩阵。
        self.wo = nn.Linear(args.n_heads * self.head_dim, args.dim, bias=False)

        # 定义dropout。
        self.attn_dropout = nn.Dropout(args.dropout)
        self.resid_dropout = nn.Dropout(args.dropout)
        # 保存dropout概率。
        self.dropout = args.dropout

        # 检查是否使用Flash Attention(需要PyTorch >= 2.0)。
        self.flash = hasattr(torch.nn.functional, 'scaled_dot_product_attention')
        if not self.flash:
            # 若不支持Flash Attention,则使用手动实现的注意力机制,并设置mask。
            print("WARNING: using slow attention. Flash Attention requires PyTorch >= 2.0")
            # 创建一个上三角矩阵,用于遮蔽未来信息。
            mask = torch.full((1, 1, args.max_seq_len, args.max_seq_len), float("-inf"))
            mask = torch.triu(mask, diagonal=1)
            # 注册为模型的缓冲区
            self.register_buffer("mask", mask)

    def forward(self, x: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor):
        # 获取批次大小和序列长度,[batch_size, seq_len, dim]
        bsz, seqlen, _ = x.shape

        # 计算查询(Q)、键(K)、值(V)。
        xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)
        # 调整形状以适应头的维度。
        xq = xq.view(bsz, seqlen, self.n_local_heads, self.head_dim)
        xk = xk.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)
        xv = xv.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)

        # 应用旋转位置嵌入(RoPE)。
        xq, xk = apply_rotary_emb(xq, xk, freqs_cos, freqs_sin)

        # 对键和值进行扩展以适应重复次数。
        xk = repeat_kv(xk, self.n_rep)
        xv = repeat_kv(xv, self.n_rep)

        # 将头作为批次维度处理。
        xq = xq.transpose(1, 2)
        xk = xk.transpose(1, 2)
        xv = xv.transpose(1, 2)

        # 根据是否支持Flash Attention,选择实现方式。
        if self.flash:
            # 使用Flash Attention。
            output = torch.nn.functional.scaled_dot_product_attention(xq, xk, xv, attn_mask=None, dropout_p=self.dropout if self.training else 0.0, is_causal=True)
        else:
            # 使用手动实现的注意力机制。
            scores = torch.matmul(xq, xk.transpose(2, 3)) / math.sqrt(self.head_dim)
            assert hasattr(self, 'mask')
            scores = scores + self.mask[:, :, :seqlen, :seqlen]
            scores = F.softmax(scores.float(), dim=-1).type_as(xq)
            scores = self.attn_dropout(scores)
            output = torch.matmul(scores, xv)

        # 恢复时间维度并合并头。
        output = output.transpose(1, 2).contiguous().view(bsz, seqlen, -1)

        # 最终投影回残差流。
        output = self.wo(output)
        output = self.resid_dropout(output)
        return output
Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐