抛砖引玉:浅谈ROPE位置编码模式下,q、k的分布(均值与方差)对注意力远程衰减的影响
October 13, 2024
1、引言
作者在论文中指出:虽然普遍认为 RoPE 的有用之处在于它有助于随着相对距离的增加而衰减 token 之间的依赖性(这部分可以参阅苏神的帖子:"Transformer升级之路:2、博采众长的旋转式位置编码"),但该论文作者认为这不太可能是主要原因。因为作者实验发现当都为均值为0的高斯初始化方案时,远程衰减性并不存在。

2、实验探查
这一下子颠覆了固有的认知,我觉得还是很有必要做下本地实验,验证下是否的确如此。
import torch import numpy as np import matplotlib.pyplot as plt device = 'cuda' seq_len = 5000 output_dim = 768 frequency = 10000 batch_size = 1 position_ids = torch.arange(0, seq_len, dtype=torch.float).unsqueeze(-1) indices = torch.arange(0, output_dim // 2, dtype=torch.float) indices = torch.pow(frequency, -2 * indices / output_dim) embeddings = position_ids * indices embeddings = torch.stack([torch.sin(embeddings), torch.cos(embeddings)], dim=-1) embeddings = embeddings.repeat((batch_size, *([1]*len(embeddings.shape)))) embeddings = torch.reshape(embeddings, (batch_size, seq_len, output_dim)) embeddings = embeddings.to(device) pos_emb = embeddings cos_pos = pos_emb[..., 1::2].repeat_interleave(2, dim=-1) sin_pos = pos_emb[..., ::2].repeat_interleave(2, dim=-1) mean = 0 # initilize the q k with Gaussian distribution torch.manual_seed(333) q = torch.normal(mean=mean, std=1.0, size=(1, seq_len, output_dim)).to(device) k = torch.normal(mean=mean, std=1.0, size=(1, seq_len, output_dim)).to(device) q2 = torch.stack([-q[..., 1::2], q[...,::2]], -1) q2 = q2.reshape(q.shape) k2 = torch.stack([-k[..., 1::2], k[...,::2]], -1) k2 = k2.reshape(k.shape) q_rope = q * cos_pos + q2 * sin_pos k_rope = k * cos_pos + k2 * sin_pos Activation_score_original = torch.einsum('bmd,bnd->bmn', q, k) Activation_score_rope = torch.einsum('bmd,bnd->bmn', q_rope, k_rope) score_decay = torch.flip(Activation_score_rope[0][-1]/torch.max(Activation_score_rope[0][-1]), dims=[0]).cpu().numpy() plt.plot(score_decay) plt.title('Activation Decay Test') plt.xlabel('Seqence_len') plt.ylabel('Activation score') plt.title(f'Mean of initial Gaussian Distribution: {mean}') plt.tight_layout() plt.show()
结果如下:

的确如该论文所言,注意力远程衰减的性质并不存在。
那问题出在哪里呢?毕竟ROPE现在广泛地应用在主流的大模型框架中,效果层面应该是接受了事实的检验的。
现在主流大模型突破长文本的优化训练方式之一就是增加(比如从10000增加到1000000)
笔者思考后,觉得可以试下对论文中Proposition 3.2部分的假设进行放松,查看注意力衰减的效果,具体如下:
2-1)放松均值为0的假设,查看不同均值情况下的注意力远程衰减效果:
means = list(np.arange(-1, 1.25, 0.25)) # Different mean values for Gaussian distribution num_subplots = len(means) num_cols = 2 num_rows = (num_subplots + 1) // 2 fig, axes = plt.subplots(num_rows, num_cols, figsize=(10, 10)) # Adjust the figsize as needed fig.suptitle('Activation Decay Test') for i, mean_val in enumerate(means): # initilize the q k with Gaussian distribution torch.manual_seed(333) q = torch.normal(mean=mean_val, std=1.0, size=(1, seq_len, output_dim)).to(device) k = torch.normal(mean=mean_val, std=1.0, size=(1, seq_len, output_dim)).to(device) q2 = torch.stack([-q[..., 1::2], q[...,::2]], -1) q2 = q2.reshape(q.shape) k2 = torch.stack([-k[..., 1::2], k[...,::2]], -1) k2 = k2.reshape(k.shape) q_rope = q * cos_pos + q2 * sin_pos k_rope = k * cos_pos + k2 * sin_pos Activation_score_original = torch.einsum('bmd,bnd->bmn', q, k) Activation_score_rope = torch.einsum('bmd,bnd->bmn', q_rope, k_rope) score_decay = torch.flip(Activation_score_rope[0][-1]/torch.max(Activation_score_rope[0][-1]) , dims=[0]).cpu().numpy() row_idx = i // num_cols col_idx = i % num_cols axes[row_idx, col_idx].plot(score_decay) axes[row_idx, col_idx].set_xlabel('Seqence_len') axes[row_idx, col_idx].set_ylabel('Activation score') axes[row_idx, col_idx].set_title(f'Mean of initial Gaussian Distribution: {mean_val}') plt.tight_layout() plt.show()
结果如下:

我们可以发现:
- 当初始化的高斯分布的均值绝对值越大时,注意力远程衰减性质越明显;当均值趋于0时,远程衰减性质消失。
上面的实验中,我们默认设置的是,即均值同向,我们再试验下当(均值异向)时的效果,具体设置如下:
torch.manual_seed(333) q = torch.normal(mean=1, std=1.0, size=(1, seq_len, output_dim)).to(device) k = torch.normal(mean=-1, std=1.0, size=(1, seq_len, output_dim)).to(device)

可以看到,反而出现了注意力远程增加的性质!
综上,我们目前得到的信息是:
在ROPE位置编码中,注意力远程衰减的性质与的分布均值关系密切;如果分布的均值同向且数值大的话,注意力远程衰减的性质越强;反之,则相反;
那么反向推理的话,我们是否可以认为是模型在训练过程中,针对不同的层、不同的注意力头,模型学到了不同的分布(对应不同的分布均值)从而实现了不同的注意力远程衰减的性质呢?
- 有些层、注意力头可能更注重远程衰减,从而更关注局部信息
- 有些层、注意力头可能基本不具备远程衰减性质,从而更关注全局信息
- 而这些是通过训练实现不同层、注意力头的分布实现的
而
其中,为该层的输入,为待训练的权重,为了更好地学习不同的分布均值,保留bias选项是否更好呢?
然后我们以Qwen2.5系列模型为例看一下所有的nn.Linear层哪些添加了bias,哪些没有呢?

可见qwen2.5系列为权重是保留了bias的,其他的则默认都没有保留bias,这么巧?
2-2)放松方差为1的假设,查看不同方差情况下的注意力远程衰减效果:
stds = list(np.arange(1, 3, 0.5)) # Different mean values for Gaussian distribution num_subplots = len(stds) num_cols = 2 num_rows = (num_subplots + 1) // 2 fig, axes = plt.subplots(num_rows, num_cols, figsize=(10, 10)) # Adjust the figsize as needed fig.suptitle('Activation Decay Test') for i, std_val in enumerate(stds): # initilize the q k with Gaussian distribution torch.manual_seed(333) q = torch.normal(mean=1, std=std_val, size=(1, seq_len, output_dim)).to(device) k = torch.normal(mean=1, std=std_val, size=(1, seq_len, output_dim)).to(device) q2 = torch.stack([-q[..., 1::2], q[...,::2]], -1) q2 = q2.reshape(q.shape) k2 = torch.stack([-k[..., 1::2], k[...,::2]], -1) k2 = k2.reshape(k.shape) q_rope = q * cos_pos + q2 * sin_pos k_rope = k * cos_pos + k2 * sin_pos Activation_score_original = torch.einsum('bmd,bnd->bmn', q, k) Activation_score_rope = torch.einsum('bmd,bnd->bmn', q_rope, k_rope) score_decay = torch.flip(Activation_score_rope[0][-1]/torch.max(Activation_score_rope[0][-1]) , dims=[0]).cpu().numpy() row_idx = i // num_cols col_idx = i % num_cols axes[row_idx, col_idx].plot(score_decay) axes[row_idx, col_idx].set_xlabel('Seqence_len') axes[row_idx, col_idx].set_ylabel('Activation score') axes[row_idx, col_idx].set_title(f'std of initial Gaussian Distribution: {std_val}') plt.tight_layout() plt.show()
结果如下:

我们可以发现:
当初始化的高斯分布的方差越大时,注意力远程衰减性质越弱;越小时,注意力远程衰减性质越好。
从这个性质出发,在参数初始化环节中设置的方差不要太大,其实也是有助于训练过程中有比较好的注意力远程衰减性质。
比如GPT-2系列参数初始化的方差设置为0.02
总结、猜想与不足
ROPE位置编码模式下,注意力远程衰减的性质与q、k的分布息息相关。当我们假定q、k分布服从高斯分布的前提下,具体表现为:
- 1)q、k均值同向(都大于0,或者都小于0)时,均值绝对值越大,注意力远程衰减性质越明显;
- 2)q、k均值异向(一个大于0,一个小于0)时,甚至出现注意力远程增加的性质;
- 3)q、k方差越大,注意力远程衰减性质越弱,越小时,注意力远程衰减性质越好;
因为经过训练后,不同层下不同注意力头生成的q、k分布各异,所以赋予了不同层、不同注意力头下具备各异的注意力远程衰减力度(有些关注局部、有些关注整体),甚至个别呈现出注意力远程增加性质(适配长上下文的插针测试任务)。
因为q、k的分布与注意力远程衰减的关系,所以初始化方案中,设置bias选项默认为True更好,初始化的方差不能太大
不足:
本文只能给出实验性的结论,暂未给出ROPE下注意力远程衰减与q、k的分布关系的理论性的推导证明;本文抛砖引玉,期待看到社区大佬们进一步的深入研究。