在自然语言处理和序列建模任务中,元素的顺序信息至关重要。然而,自注意力机制本身并不包含序列顺序信息,因为它是基于元素之间的相似性计算的,与它们在序列中的位置无关。
问题的数学表述:
给定两个序列:
自注意力机制会得到相同的相似度矩阵,因为元素对之间的相似性是相同的,但这两个序列的含义完全不同。
位置编码的作用:
位置编码为每个位置i生成一个向量\mathbf{p}_i \in \mathbb{R}^{d_{\text{model}}},将其与输入向量\mathbf{x}_i相加:
这样,模型就能够感知到元素之间的相对位置关系。
可学习性:
周期性:
可扩展性:
区分性:
原始正弦位置编码:
Vaswani等人在《Attention Is All You Need》中提出的正弦位置编码定义为:
对于位置i和维度j:
其中:
数学性质:
相对位置的重要性:
在许多序列处理任务中,相对位置比绝对位置更重要。例如:
相对位置编码的数学表示:
import torch import torch.nn as nn import math class SinusoidalPositionalEncoding(nn.Module): """ 正弦位置编码实现 """ def __init__(self, d_model, max_len=5000): super().__init__() self.d_model = d_model self.max_len = max_len # 创建位置编码矩阵 position = torch.arange(max_len, dtype=torch.float).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe = torch.zeros(max_len, d_model) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) self.register_buffer('pe', pe) def forward(self, x): """ Args: x: [batch_size, seq_len, d_model] """ # 将位置编码添加到输入中 return x + self.pe[:x.size(1)] class RelativePositionalEncoding(nn.Module): """ 相对位置编码实现 """ def __init__(self, d_model, max_relative_position=50): super().__init__() self.d_model = d_model self.max_relative_position = max_relative_position # 创建相对位置编码表 self.relative_positions = nn.Parameter( torch.zeros(2 * max_relative_position + 1, d_model) ) # 初始化相对位置编码 nn.init.xavier_uniform_(self.relative_positions) def forward(self, seq_len): """ Args: seq_len: 序列长度 Returns: relative_positions: [seq_len, seq_len, d_model] """ # 创建相对位置矩阵 positions = torch.arange(seq_len, dtype=torch.long) relative_positions = positions.unsqueeze(1) - positions.unsqueeze(0) # 将相对位置限制在范围内 relative_positions = torch.clamp( relative_positions, -self.max_relative_position, self.max_relative_position ) # 将相对位置映射到编码表 relative_positions = relative_positions + self.max_relative_position encoding = self.relative_positions[relative_positions] return encoding
学习式正弦位置编码:
其中a, b, c, d是可学习的参数。
旋转位置编码(RoPE):
近年来提出的旋转位置编码通过旋转矩阵来实现位置编码:
其中R(i)是旋转矩阵,通过复数运算实现。
优点:
缺点:
Embedding方式:
最简单的可学习位置编码是使用Embedding层:
class LearnedPositionalEncoding(nn.Module): """ 可学习位置编码 """ def __init__(self, d_model, max_len=5000): super().__init__() self.d_model = d_model self.max_len = max_len # 位置编码Embedding self.pos_embedding = nn.Embedding(max_len, d_model) self.pos_embedding.weight.data = self._init_embedding() def _init_embedding(self): """ 初始化位置编码 """ # 使用正弦初始化 pe = torch.zeros(self.max_len, self.d_model) position = torch.arange(self.max_len, dtype=torch.float).unsqueeze(1) div_term = torch.exp(torch.arange(0, self.d_model, 2).float() * (-math.log(10000.0) / self.d_model)) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) return pe def forward(self, x): """ Args: x: [batch_size, seq_len, d_model] """ batch_size, seq_len, d_model = x.shape positions = torch.arange(seq_len, device=x.device).unsqueeze(0).repeat(batch_size, 1) pos_encoding = self.pos_embedding(positions) return x + pos_encoding
结合正弦和学习式:
class HybridPositionalEncoding(nn.Module): """ 混合位置编码:结合正弦和学习式 """ def __init__(self, d_model, max_len=5000, alpha=0.1): super().__init__() self.d_model = d_model self.max_len = max_len self.alpha = alpha # 固定正弦编码 self.sinusoidal_pe = SinusoidalPositionalEncoding(d_model, max_len) # 可学习编码 self.learned_pe = LearnedPositionalEncoding(d_model, max_len) # 学习权重 self.alpha = nn.Parameter(torch.tensor(alpha)) def forward(self, x): """ Args: x: [batch_size, seq_len, d_model] """ sin_pe = self.sinusoidal_pe(x) learned_pe = self.learned_pe(x) return self.alpha * learned_pe + (1 - self.alpha) * sin_pe
根据任务选择:
根据资源选择:
性能比较:
在不同的基准任务上测试不同位置编码的表现:
| 编码方法 | WMT14英德 | WMT14英法 | 翻译速度 | 内存占用 |
|---|---|---|---|---|
| 正弦编码 | 28.4 | 41.8 | 快 | 低 |
| 学习编码 | 29.1 | 42.3 | 中 | 中 |
| 混合编码 | 29.3 | 42.5 | 慢 | 高 |
| RoPE | 29.2 | 42.4 | 快 | 低 |
应用建议:
基本原理:
将维度分成多个组,每组使用不同的位置编码策略:
class GroupPositionalEncoding(nn.Module): """ 分组位置编码 """ def __init__(self, d_model, num_groups=4, encoding_type='sinusoidal'): super().__init__() self.d_model = d_model self.num_groups = num_groups self.group_size = d_model // num_groups # 每组使用不同的位置编码 self.encodings = nn.ModuleList() for _ in range(num_groups): if encoding_type == 'sinusoidal': self.encodings.append( SinusoidalPositionalEncoding(self.group_size) ) elif encoding_type == 'learned': self.encodings.append( LearnedPositionalEncoding(self.group_size) ) def forward(self, x): """ Args: x: [batch_size, seq_len, d_model] """ batch_size, seq_len, d_model = x.shape output = torch.zeros_like(x) # 按组处理 for i in range(self.num_groups): start_idx = i * self.group_size end_idx = (i + 1) * self.group_size group_x = x[:, :, start_idx:end_idx] output[:, :, start_idx:end_idx] = self.encodings[i](group_x) return output
基于序列长度的调整:
class AdaptivePositionalEncoding(nn.Module): """ 自适应位置编码:根据序列长度动态调整编码 """ def __init__(self, d_model, max_len=5000): super().__init__() self.d_model = d_model self.max_len = max_len # 计算最优的编码参数 self.encoding_params = nn.ParameterList([ nn.Parameter(torch.randn(d_model // 2)) for _ in range(3) # 长度、中、短 ]) def forward(self, x): """ Args: x: [batch_size, seq_len, d_model] """ seq_len = x.shape[1] # 根据序列长度选择编码参数 if seq_len < 100: params = self.encoding_params[0] # 短序列 elif seq_len < 1000: params = self.encoding_params[1] # 中等序列 else: params = self.encoding_params[2] # 长序列 # 生成位置编码 position = torch.arange(seq_len, device=x.device).unsqueeze(1) encoding = torch.zeros(seq_len, self.d_model, device=x.device) encoding[:, 0::2] = torch.sin(position / (10000 ** (-params.unsqueeze(1)))) encoding[:, 1::2] = torch.cos(position / (10000 ** (-params.unsqueeze(1)))) return x + encoding
分类位置编码:
除了加法操作,还可以使用拼接方式:
class ConcatPositionalEncoding(nn.Module): """ 拼接位置编码 """ def __init__(self, d_model, pos_dim=128): super().__init__() self.d_model = d_model self.pos_dim = pos_dim # 位置编码投影 self.pos_projection = nn.Linear(d_model + pos_dim, d_model) # 位置编码生成 self.pos_encoding = SinusoidalPositionalEncoding(pos_dim) def forward(self, x): """ Args: x: [batch_size, seq_len, d_model] """ pos_encoding = self.pos_encoding(x)[:, :, :self.pos_dim] # 拼接位置编码 combined = torch.cat([x, pos_encoding], dim=-1) # 投影回原维度 return self.pos_projection(combined)
调试技巧:
常见问题:
优化建议: