以下代码实现了旋转位置编码(RoPE),请阅读代码并填写空缺部分。 [1]和[2]处应分别填入:
import torch
import math
def apply_rotary_pos_emb(q, k, pos):
"""
对 Query 和 Key 应用旋转位置编码(RoPE)
参数:
q: Query 张量, shape(batch, heads, seq_len, head_dim)
k: Key 张量, shape(batch, heads, seq_len, head_dim)
pos: 位置索引, shape(seq_len,)
返回:
旋转编码后的 q 和 k
"""
head_dim = q.size(-1)
# 计算频率基底:theta_i = 1 / (10000 ^ (2i/d))
freq_indices = torch.arange(0, head_dim, 2, dtype=torch.float32)
freqs = 1.0 / (10000.0 ** (freq_indices / head_dim))
# 计算位置角度:pos * theta
angles = pos.unsqueeze(-1).float() * freqs.unsqueeze(0) # (seq_len, head_dim/2)
cos_vals = torch.cos(angles) # (seq_len, head_dim/2)
sin_vals = torch.sin(angles) # (seq_len, head_dim/2)
# 将 q 拆分为偶数维和奇数维
q_even = q[..., 0::2] # (..., head_dim/2)
q_odd = q[..., 1::2] # (..., head_dim/2)
# [1] 对 Query 应用旋转变换:q_rotated_even = q_even * cos - q_odd * sin
________________________________________________________________________________
# [2] 对 Query 应用旋转变换:q_rotated_odd = q_even * sin + q_odd * cos
________________________________________________________________________________
# 将偶数维和奇数维交错合并
q_rotated = torch.stack([q_rotated_even, q_rotated_odd], dim=-1)
q_rotated = q_rotated.flatten(start_dim=-2)
# 对 k 应用相同的旋转变换
k_even = k[..., 0::2]
k_odd = k[..., 1::2]
k_rotated_even = k_even * cos_vals - k_odd * sin_vals
k_rotated_odd = k_even * sin_vals + k_odd * cos_vals
k_rotated = torch.stack([k_rotated_even, k_rotated_odd], dim=-1)
k_rotated = k_rotated.flatten(start_dim=-2)
return q_rotated, k_rotated
[1] q_rotated_even = q_even * cos_vals - q_odd * sin_vals
[2] q_rotated_odd = q_even * sin_vals + q_odd * cos_vals
[1] q_rotated_even = q_even * sin_vals - q_odd * cos_vals
[2] q_rotated_odd = q_even * cos_vals + q_odd * sin_vals
[1] q_rotated_even = q_even + cos_vals
[2] q_rotated_odd = q_odd + sin_vals
[1] q_rotated_even = q_even * cos_vals + q_odd * sin_vals
[2] q_rotated_odd = q_even * sin_vals - q_odd * cos_vals