【大模型】RoPE原理与实践中的一些思考(附代码)~
RoPE原理
不讲从0-1的推导原理
背景是Transformer的positional 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} word−embedding
[12]+positional−embedding
[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θ)=real−after−rotateθ (acosθ−bsinθ)+iimaginary−after−rotateθ (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×d′∗f(freq
1⋅10000−d2⋅1⋯seq⋅10000−d2⋅1⋯cos(mθk)⋯1⋅10000−d2⋅d′⋯seq⋅10000−d2⋅d′
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} qr′qi′=qr∗cos(freq)−qi∗sin(freq)=qr∗sin(freq)+qi∗cos(freq)
苏神在博客中写道
[cosmθ0−sinmθ000⋯00sinmθ0cosmθ000⋯0000cosmθ1−sinmθ1⋯0000sinmθ1cosmθ1⋯00⋯⋯⋯⋯⋯⋯⋯0000⋯cosmθd/2−1−sinmθd/2−10000⋯sinmθd/2−1cosmθ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θ000⋯00−sinmθ0cosmθ000⋯0000cosmθ1sinmθ1⋯0000−sinmθ1cosmθ1⋯00⋯⋯⋯⋯⋯⋯⋯0000⋯cosmθd/2−1sinmθd/2−10000⋯−sinmθd/2−1cosmθd/2−1
q0q1q2q3⋯qd−2qd−1
并且定义
[q0′q1′]=[cosmθ0−sinmθ0sinmθ0cosmθ0][q0q1][q2′q3′]=[cosmθ1−sinmθ1sinmθ1cosmθ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} [q0′q1′][q2′q3′]=[cosmθ0sinmθ0−sinmθ0cosmθ0][q0q1]=[cosmθ1sinmθ1−sinmθ1cosmθ1][q2q3]⋯
去实现q=[q0 q1 ... q_{d-1}] -> q'=[q'0 q'1 ... q'{d-1}]的旋转
[!important]
这里容易误解的点在于,苏神是对f(q,m)=[cosmθ−sinmθsinmθcosmθ][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,⋯,qd−1]
因此不要忘记要对除
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}q0∈Rbs×1×n−head×1,切记这里的q0q_0q0实则是q(m,0)q^{(m,0)}q(m,0),和它交互的是cosmθ0\cos m\theta_0cosmθ0与sinmθ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×n−head×1,与它交互的是cosmθk\cos m\theta_kcosmθk与sinmθk\sin m\theta_ksinmθk。从元素级别可以看到每一个元素的交互都是一样的。
[!important]
事实上,拓展到高维后,有结论如下:
对于任意维度ddd,任意第kkk组k=0,1,...,d/2−1k=0,1,...,d/2-1k=0,1,...,d/2−1:
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}q∈Rseq×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(seq−1,0)q(0,1)q(1,1)⋯q(seq−1,1)⋯⋯⋯⋯q(0,d−1)q(1,d−1)⋯q(seq−1,d−1)
根据上述结论可知:
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(seq−1,0)q(0,2)q(1,2)⋯q(seq−1,2)⋯⋯⋯⋯q(0,d−2)q(1,d−2)⋯q(seq−1,d−2) = q(0,1)q(1,1)⋯q(seq−1,1)q(0,3)q(1,3)⋯q(seq−1,3)⋯⋯⋯⋯q(0,d−1)q(1,d−1)⋯q(seq−1,d−1)
写到这里,也更加清晰,苏神的推导中取q\boldsymbol{q}q逐mmm行进行RmR_mRm矩阵旋转:
q′(m,2i)=q(m,2i)⋅cosmθi−q(m,2i+1)⋅sinmθiq′(m,2i+1)=q(m,2i)⋅sinmθi+q(m,2i+1)⋅cosmθ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θi−q(m,2i+1)⋅sinmθi=q(m,2i)⋅sinmθi+q(m,2i+1)⋅cosmθi
而工程代码中总结规律发现上述方法效率太低,既然:
- q(m,2i)q^{(m,2i)}q(m,2i)和q(m,2i+1)q^{(m,2i+1)}q(m,2i+1)总是有固定的cosmθi\cos m\theta_icosmθi/sinmθi\sin m\theta_isinmθi和−sinmθi-\sin m\theta_i−sinmθi/cosmθi\cos m\theta_icosmθi要相乘;
- 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}qr和qi\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])⏟cosmθ−[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])⏟sinmθ[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])⏟sinmθ+[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])⏟cosmθ \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′(seq−1,0)q′(0,2)q′(1,2)⋯q′(seq−1,2)⋯⋯⋯⋯q′(0,d−2)q′(1,d−2)⋯q′(seq−1,d−2) qi′ q′(0,1)q′(1,1)⋯q′(seq−1,1)q′(0,3)q′(1,3)⋯q′(seq−1,3)⋯⋯⋯⋯q′(0,d−1)q′(1,d−1)⋯q′(seq−1,d−1) =qr q(0,0)q(1,0)⋯q(seq−1,0)q(0,2)q(1,2)⋯q(seq−1,2)⋯⋯⋯⋯q(0,d−2)q(1,d−2)⋯q(seq−1,d−2) ∗cosmθ cos( 0θ01θ0⋯(seq−1)θ00θ11θ1⋯(seq−1)θ1⋯⋯⋯⋯0θd/2−11θd/2−1⋯(seq−1)θd/2−1 )−qi q(0,1)q(1,1)⋯q(seq−1,1)q(0,3)q(1,3)⋯q(seq−1,3)⋯⋯⋯⋯q(0,d−1)q(1,d−1)⋯q(seq−1,d−1) ∗sinmθ sin( 0θ01θ0⋯(seq−1)θ00θ11θ1⋯(seq−1)θ1⋯⋯⋯⋯0θd/2−11θd/2−1⋯(seq−1)θd/2−1 )=qr q(0,0)q(1,0)⋯q(seq−1,0)q(0,2)q(1,2)⋯q(seq−1,2)⋯⋯⋯⋯q(0,d−2)q(1,d−2)⋯q(seq−1,d−2) ∗sinmθ sin( 0θ01θ0⋯(seq−1)θ00θ11θ1⋯(seq−1)θ1⋯⋯⋯⋯0θd/2−11θd/2−1⋯(seq−1)θd/2−1 )+qi q(0,1)q(1,1)⋯q(seq−1,1)q(0,3)q(1,3)⋯q(seq−1,3)⋯⋯⋯⋯q(0,d−1)q(1,d−1)⋯q(seq−1,d−1) ∗cosmθ cos( 0θ01θ0⋯(seq−1)θ00θ11θ1⋯(seq−1)θ1⋯⋯⋯⋯0θd/2−11θd/2−1⋯(seq−1)θd/2−1 )
具体复杂度对比这里就不做了,工程代码实现了从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
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)