我要提问
ARTICLE DETAIL

资讯详情

前沿编程新知与开发实战干货的深度解读。

一文讲透ROPE

一文讲透ROPE 背景是基于什么样的目的出发提出的ROPE呢苏剑林提出的精妙设想能不能在输入端q和k注入绝对位置信息但是当它们进行内积时结果就恰好只和它们的相对位置有关即要找到一个函数,使得其中m和n是绝对位置序号而内积的结果最后只跟相对位置 m - n有关。怎么设计函数为了找到函数f RoPE巧妙的借助了复平面的几何性质。我们也先来回顾一下1.回顾复平面如果小伙伴对这块知识比较熟悉可以直接跳过复数通用形式a实部对应平面 X 轴坐标b虚部对应平面 Y 轴坐标i虚数单位对应这平面坐标a, b任意复数对应平面向量向量长度模对向量用极坐标表示可以表示为其中R是模长是角度。且有欧拉公式为什么可以这么表示呢向量是有大小有方向的R是大小那么如果表示的是方向是不是就OK了其实可以看做是一个单位圆可以从导数角度来理解下这个几何意义在时的导数变化方向是向着虚轴i当导数方向是负实轴这两点都是严格垂直于点和圆心的线大致可以推断就是在单位圆上运动的点这个推断理由不严谨哈仅仅是帮助大家理解的几何意义重要特性两个向量q和k的内积等于q乘以k的共轭复数的实部证明例如向量和对应(a,b)和(c,d)那么(a,b)和(c,d)的内积acdb,而那么用极坐标来表示那么向量z和y的内积2.q和k的内积由于两个向量q和k的内积等于q乘以k的共轭复数的实部则等等既然两个向量的内积是两者的角度差而RoPE最初的设想是q和k做内积之后结果之和相对位置有关。那么我们是不是只要把Q和K的序列位置m,n当做是旋转角度加进去就可以了大功告成。q在位置m变成k在位置变成此时再计算他们的内积推导大功告成我们可以清晰地看到原本各自带着绝对位置和的向量q和k一做内积公式里只剩下了相对位置。这就是RoPE被称为“旋转”的原因。实际应用具体怎么做真实模型中的特征维度d通常很大如 4096 维怎么把 2 维的复数旋转推广到高维空间1.扩展到高维空间RoPE 的做法是两两分组。把d维的向量分成d/2个 2 维平面。对每一个 2 维平面应用上述的旋转操作。为了让模型捕捉到不同尺度的位置关系每个平面的旋转频率也就是转的速度是不一样的。具体来说第组 2 维平面的旋转角频率定义为靠前的维度较小大转得快负责捕捉近距离的高频细节位置靠后的维度较大小转得慢负责捕捉远距离的低频全局位置。因为旋转的慢会带来一个极大的问题需要外推这后面再讲。2.矩阵表示 vs 实际运算优化在数学上把一个 2 维向量旋转,相当于乘以一个旋转矩阵为什么上面说过代表着角度为在单位圆上的点那么一个向量旋转就相当于乘上即那么旋转后即旋转角度矩阵因此对向量做旋转等于向量和旋转矩阵的相乘也可以写成哈达玛积逐元素相乘在原文中采用相邻2个元素为一组 的方式进行旋转但是在实际工程实现中出于效率上考虑一般不会这样做而是前后对半分再分组的方式。例如输入向量q的维度是d切成两半则分组为把整个向量合并起来对应关系为运算结果为什么这么分组呢按照原文的分组方式的话需要频繁的交错采样取奇数位和偶数位这在计算时gpu计算单元到显存中去取值的效率就会非常低。而采用前后对半分的方式元素在显存中的地址是连续的可以提升带宽利用率。3.旋转后的内积上文说过向量的内积和两者的角度差有关因此把序列向量位置m当做是旋转角度加入到向量角度中去让其内积后只和相对位置有关现在我们就来看一下旋转后的内积则上面的公式中代表逆时针旋转那这个转置又表示什么呢令那么因为偶函数f(-x)f(x)奇函数f(-x)-f(x)而 cos是偶函数而sin是奇函数那么即从直观上先旋转再旋转,是不是等于总共旋转了实际上带入矩阵实际计算的话也确实是等于上述结果的这里就省略计算过程了这里的qk是列向量所以是对q进行转置即只和n-m的相对位置有关通过绝对位置的旋转实现了相对位置的编码并且不像正余弦位置编码那样不会引入交叉噪声。实际代码实现import torch import torch.nn as nn def rotate_half(x: torch.Tensor) - torch.Tensor: 【数据变化详解】 输入 x 形状: [1, 1, 2, 4] 例取 t1 位置的向量: [q1, q2, q3, q4] 1. 前半部分切片 x1: 索引: [..., :2] 形状: [1, 1, 2, 2] 数值: [q1, q2] 2. 后半部分切片 x2: 索引: [..., 2:] 形状: [1, 1, 2, 2] 数值: [q3, q4] 3. 取负并拼接 torch.cat((-x2, x1), dim-1): -x2 数值: [-q3, -q4] 拼接后形状: [1, 1, 2, 4] 最终数值: [-q3, -q4, q1, q2] # x.shape[-1] 4, 4 // 2 2 x1 x[..., : x.shape[-1] // 2] # [1, 1, 2, 2] x2 x[..., x.shape[-1] // 2 :] # [1, 1, 2, 2] # 拼接产生的张量在物理显存中直接连续存储 out torch.cat((-x2, x1), dim-1) # [1, 1, 2, 4] return out def apply_rotary_pos_emb( q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, unsqueeze_dim: int 1, ) - tuple[torch.Tensor, torch.Tensor]: 【数据变化详解】以 Query 张量 q 为例 输入 q 形状: [1, 1, 2, 4] (Batch1, Heads1, Seq2, Dim4) 输入 cos, sin 形状: [2, 4] 1. unsqueeze 维度扩充: cos.unsqueeze(1) 形状变更为: [1, 1, 2, 4] 数据在 dim0 和 dim1 自动广播对齐 2. 逐元素计算项 1: q * cos 位置 t1 处: [q1, q2, q3, q4] * [cos(θ1), cos(θ2), cos(θ1), cos(θ2)] [q1*cos(θ1), q2*cos(θ2), q3*cos(θ1), q4*cos(θ2)] 3. 逐元素计算项 2: rotate_half(q) * sin 位置 t1 处: [-q3, -q4, q1, q2] * [sin(θ1), sin(θ2), sin(θ1), sin(θ2)] [-q3*sin(θ1), -q4*sin(θ2), q1*sin(θ1), q2*sin(θ2)] 4. 两项相加得到最终 q_embed: 第 0 维: q1*cos(θ1) - q3*sin(θ1) -- 二维平面 1 旋转后坐标 x 第 1 维: q2*cos(θ2) - q4*sin(θ2) -- 二维平面 2 旋转后坐标 x 第 2 维: q3*cos(θ1) q1*sin(θ1) -- 二维平面 1 旋转后坐标 y 第 3 维: q4*cos(θ2) q2*sin(θ2) -- 二维平面 2 旋转后坐标 y # [2, 4] - [1, 1, 2, 4] cos cos.unsqueeze(unsqueeze_dim) sin sin.unsqueeze(unsqueeze_dim) # 两个 [1, 1, 2, 4] 的张量相加结果形状保持 [1, 1, 2, 4] q_embed (q * cos) (rotate_half(q) * sin) k_embed (k * cos) (rotate_half(k) * sin) return q_embed, k_embed class LlamaRotaryEmbedding(nn.Module): def __init__( self, dim: int, max_position_embeddings: int 2048, base: int 10000, deviceNone, ): super().__init__() self.dim dim # dim 4 self.max_position_embeddings max_position_embeddings self.base base 【数据变化详解】计算逆角频率向量 inv_freq: 1. torch.arange(0, 4, 2) - 生成序列 [0, 2]长度为 dim/2 2 2. [0, 2] / 4 - [0.0, 0.5] 3. 10000 ** [0.0, 0.5] - [10000^0, 10000^0.5] [1.0, 100.0] 4. 倒数倒置 - [1.0 / 1.0, 1.0 / 100.0] [1.0, 0.01] 最终 inv_freq 形状: [2]数值: [1.0, 0.01] (对应频率 θ11.0, θ20.01) inv_freq 1.0 / ( self.base ** ( torch.arange(0, self.dim, 2, dtypetorch.int64).float().to(device) / self.dim ) ) self.register_buffer(inv_freq, inv_freq, persistentFalse) self._set_cos_sin_cache( seq_lenmax_position_embeddings, devicedevice, dtypetorch.get_default_dtype(), ) def _set_cos_sin_cache(self, seq_len: int, deviceNone, dtypetorch.float32): self.max_seq_len_cached seq_len 【数据变化详解】生成角度矩阵并构造 cos/sin 缓存表 (以 seq_len2 为例): 1. 位置向量 t: torch.arange(2) - [0, 1]形状: [2] 2. 外积计算 freqs outer(t, inv_freq): t: [0, 1] (形状 [2]) inv_freq: [1.0, 0.01] (形状 [2]) 外积矩阵 freqs 形状: [2, 2] 矩阵数值: 行 t0: [0 * 1.0, 0 * 0.01] [0.0, 0.0] 行 t1: [1 * 1.0, 1 * 0.01] [1.0, 0.01] 3. 拼接扩展 emb cat((freqs, freqs), dim-1): 拼接前: [2, 2] 拼接后 形状: [2, 4] 矩阵数值: 行 t0: [0.0, 0.0, 0.0, 0.0] 行 t1: [1.0, 0.01, 1.0, 0.01] -- 后半部分角度与前半部分复制对齐 4. 计算 cos 和 sin 表: cos_cached 形状: [2, 4] 行 t0: [cos(0), cos(0), cos(0), cos(0)] [1.0, 1.0, 1.0, 1.0] 行 t1: [cos(1.0), cos(0.01), cos(1.0), cos(0.01)] [0.540, 0.999, 0.540, 0.999] sin_cached 形状: [2, 4] 行 t0: [sin(0), sin(0), sin(0), sin(0)] [0.0, 0.0, 0.0, 0.0] 行 t1: [sin(1.0), sin(0.01), sin(1.0), sin(0.01)] [0.841, 0.009, 0.841, 0.009] t torch.arange( self.max_seq_len_cached, devicedevice, dtypetorch.int64 ).type_as(self.inv_freq) # [2] 外积 [2] - [2, 2] freqs torch.outer(t, self.inv_freq) # [2, 2] cat [2, 2] - [2, 4] emb torch.cat((freqs, freqs), dim-1) self.register_buffer(cos_cached, emb.cos().to(dtype), persistentFalse) self.register_buffer(sin_cached, emb.sin().to(dtype), persistentFalse) def forward( self, x: torch.Tensor, seq_len: int None ) - tuple[torch.Tensor, torch.Tensor]: if seq_len self.max_seq_len_cached: self._set_cos_sin_cache(seq_lenseq_len, devicex.device, dtypex.dtype) # 切片提取前 seq_len 行数据形状: [seq_len, dim] - [2, 4] return ( self.cos_cached[:seq_len].to(dtypex.dtype), self.sin_cached[:seq_len].to(dtypex.dtype), ) # # 步步跟踪验证脚本 # if __name__ __main__: # 使用全 1/固定数值的输入向量便于直观观察数值计算过程 batch_size 1 num_heads 1 seq_len 2 head_dim 4 print( 1. 初始化简单测试向量 ) # 构建特定的 q 矩阵便于跟踪计算 # t0: [1.0, 1.0, 1.0, 1.0] # t1: [2.0, 2.0, 2.0, 2.0] q torch.tensor( [[[[1.0, 1.0, 1.0, 1.0], [2.0, 2.0, 2.0, 2.0]]]], dtypetorch.float32 ) print(f输入 q 形状: {q.shape}) print(ft1 位置输入向量: {q[0, 0, 1].numpy()}\n) print( 2. 执行 rotate_half(q) 数据变换 ) q_rev rotate_half(q) print(frotate_half(q) 形状: {q_rev.shape}) print(ft1 原始向量 [q1, q2, q3, q4]: {q[0, 0, 1].numpy()}) print(ft1 转换向量 [-q3, -q4, q1, q2]: {q_rev[0, 0, 1].numpy()}\n) print( 3. 旋转矩阵 (cos/sin) 生成 ) rope LlamaRotaryEmbedding(dimhead_dim, max_position_embeddings10) cos, sin rope(q, seq_lenseq_len) print(fcos 缓存矩阵 [seq_len2, dim4]:\n{cos.numpy().round(4)}) print(fsin 缓存矩阵 [seq_len2, dim4]:\n{sin.numpy().round(4)}\n) print( 4. 旋转编码施加过程 (以 t1 为例) ) q_embed, _ apply_rotary_pos_emb(q, q, cos, sin) # 手动用公式推导 t1 的预期结果: # cos_row [cos(1.0), cos(0.01), cos(1.0), cos(0.01)] [0.5403, 0.9999, 0.5403, 0.9999] # sin_row [sin(1.0), sin(0.01), sin(1.0), sin(0.01)] [0.8415, 0.0100, 0.8415, 0.0100] # q_t1 [2.0, 2.0, 2.0, 2.0] # q_rev_t1 [-2.0, -2.0, 2.0, 2.0] # q1 2.0 * 0.5403 (-2.0) * 0.8415 1.0806 - 1.6830 -0.6024 # q2 2.0 * 0.9999 (-2.0) * 0.0100 1.9998 - 0.0200 1.9798 # q3 2.0 * 0.5403 2.0 * 0.8415 1.0806 1.6830 2.7636 # q4 2.0 * 0.9999 2.0 * 0.0100 1.9998 0.0200 2.0198 print( 代码算出的 t1 旋转后向量:\n, q_embed[0, 0, 1].detach().numpy().round(4), ) print(手动推导的 t1 理论结果向量:\n [-0.6024 1.9798 2.7636 2.0198])
返回列表