别再死记硬背Transformer公式了!手把手带你用PyTorch从零实现位置编码与Mask机制
从零实现Transformer核心组件:位置编码与Mask机制的PyTorch实战指南
当你第一次翻开Transformer论文时,那些复杂的数学公式是否让你望而却步?作为自然语言处理领域的革命性架构,Transformer的核心思想其实可以通过代码直观理解。本文将带你用PyTorch从零实现两个关键组件——位置编码和Mask机制,通过编写可运行的代码来反向理解其设计原理。
1. 位置编码:让模型记住词序的魔法
传统RNN天然具有处理序列的能力,而Transformer作为无时序的架构,需要额外机制来编码位置信息。这就是位置编码(Positional Encoding)的用武之地。
1.1 为什么需要位置编码?
想象你在阅读这句话:"猫追老鼠"和"老鼠追猫"。词序不同,语义完全相反。Transformer的自注意力机制本身无法感知词序,因此需要显式地注入位置信息。
位置编码的设计需要满足几个关键特性:
- 唯一性 :每个位置有唯一编码
- 相对位置关系 :能够表示位置间的相对距离
- 泛化性 :能处理比训练时更长的序列
1.2 正弦余弦编码的数学实现
原始论文使用交替的正弦和余弦函数来生成位置编码。对于位置pos和维度i,编码公式为:
PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
这种设计有几个精妙之处:
- 正弦函数的 周期性 让模型能学习到相对位置关系
- 不同维度使用不同的 波长 ,形成多层次的位置感知
- 数值范围被控制在[-1,1]之间,与词嵌入尺度匹配
1.3 PyTorch完整实现
让我们用PyTorch实现这一机制:
import torch
import math
class PositionalEncoding(torch.nn.Module):
def __init__(self, d_model: int, max_len: int = 5000):
super().__init__()
# 创建位置编码矩阵
position = torch.arange(max_len).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
pe = torch.zeros(max_len, d_model)
pe[:, 0::2] = torch.sin(position * div_term) # 偶数维度用sin
pe[:, 1::2] = torch.cos(position * div_term) # 奇数维度用cos
self.register_buffer('pe', pe) # 不参与训练
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x: Tensor, shape [batch_size, seq_len, embedding_dim]
"""
x = x + self.pe[:x.size(1)] # 只取前seq_len个位置
return x
这段代码的几个关键点:
- 使用矩阵运算而非循环,大幅提升效率
register_buffer确保位置编码不参与训练- 直接加到词嵌入上,实现位置信息注入
提示:实际应用中,位置编码通常与词嵌入相加而非拼接,这能减少模型参数同时保持信息融合
1.4 位置编码的可视化分析
让我们观察生成的位置编码矩阵:
import matplotlib.pyplot as plt
d_model = 512
max_len = 100
pe = PositionalEncoding(d_model, max_len).pe
plt.figure(figsize=(12, 6))
plt.imshow(pe[:100, :100].T, cmap='viridis')
plt.xlabel('Position')
plt.ylabel('Dimension')
plt.colorbar()
plt.show()
你会看到:
- 低维度(图底部)变化剧烈,捕获细粒度位置信息
- 高维度(图顶部)变化平缓,捕获宏观位置关系
- 相邻位置编码相似但有差异,形成平滑过渡
2. Mask机制:控制信息流的阀门
Transformer中有两种关键Mask:填充Mask(pad_mask)和序列Mask(tril_mask)。它们分别解决不同问题。
2.1 填充Mask:处理变长序列
在实际应用中,我们通常将多个句子打包成批次(batch)处理。由于句子长度不一,需要填充(pad)到相同长度。填充Mask确保模型忽略这些无意义的填充位置。
实现原理 :
- 标记所有填充位置为True
- 将这些位置对应的注意力权重设为负无穷
- 经过softmax后,这些位置的权重变为0
def create_pad_mask(seq: torch.Tensor, pad_idx: int) -> torch.Tensor:
"""
创建填充mask
Args:
seq: 输入序列,形状 [batch_size, seq_len]
pad_idx: 填充token的索引
Returns:
mask: 形状 [batch_size, 1, 1, seq_len]
"""
mask = (seq == pad_idx).unsqueeze(1).unsqueeze(2) # [batch_size, 1, 1, seq_len]
return mask
2.2 序列Mask:防止信息泄露
在解码器中,预测下一个词时不应看到"未来"的信息。序列Mask(又称因果Mask)通过上三角矩阵实现这一点。
def create_seq_mask(seq_len: int) -> torch.Tensor:
"""
创建序列mask(上三角矩阵)
Args:
seq_len: 序列长度
Returns:
mask: 形状 [1, seq_len, seq_len]
"""
mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1).bool()
return mask.unsqueeze(0) # [1, seq_len, seq_len]
注意:在训练解码器时,需要同时使用填充Mask和序列Mask,通过逻辑或操作合并它们
2.3 组合Mask的实现
实际应用中,我们经常需要组合多种Mask:
def combine_masks(pad_mask: torch.Tensor, seq_mask: torch.Tensor) -> torch.Tensor:
"""
组合填充mask和序列mask
Args:
pad_mask: 形状 [batch_size, 1, 1, seq_len]
seq_mask: 形状 [1, seq_len, seq_len]
Returns:
组合后的mask: 形状 [batch_size, 1, seq_len, seq_len]
"""
if pad_mask is None and seq_mask is None:
return None
combined_mask = pad_mask | seq_mask if pad_mask is not None and seq_mask is not None else \
pad_mask if pad_mask is not None else seq_mask
return combined_mask
3. 注意力机制中的Mask应用
理解了Mask的原理后,我们来看如何在注意力机制中应用它们。
3.1 缩放点积注意力实现
def scaled_dot_product_attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor = None
) -> torch.Tensor:
"""
缩放点积注意力计算
Args:
q: query, 形状 [batch_size, n_heads, seq_len, d_k]
k: key, 形状 [batch_size, n_heads, seq_len, d_k]
v: value, 形状 [batch_size, n_heads, seq_len, d_v]
mask: 可选的mask, 形状 [batch_size, 1, seq_len, seq_len]
Returns:
注意力输出: 形状 [batch_size, n_heads, seq_len, d_v]
"""
d_k = q.size(-1)
attn_scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k)
if mask is not None:
attn_scores = attn_scores.masked_fill(mask == 1, float('-inf'))
attn_weights = torch.softmax(attn_scores, dim=-1)
output = torch.matmul(attn_weights, v)
return output
3.2 多头注意力中的Mask处理
在完整的Transformer实现中,Mask需要适配多头注意力:
class MultiHeadAttention(torch.nn.Module):
def __init__(self, d_model: int, n_heads: int):
super().__init__()
assert d_model % n_heads == 0, "d_model必须能被n_head整除"
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads
self.w_q = torch.nn.Linear(d_model, d_model)
self.w_k = torch.nn.Linear(d_model, d_model)
self.w_v = torch.nn.Linear(d_model, d_model)
self.w_o = torch.nn.Linear(d_model, d_model)
def forward(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor = None
) -> torch.Tensor:
batch_size = q.size(0)
# 线性投影
q = self.w_q(q).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
k = self.w_k(k).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
v = self.w_v(v).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
# 应用mask(如果需要)
if mask is not None:
mask = mask.unsqueeze(1) # 适配多头 [batch_size, 1, 1, seq_len] -> [batch_size, 1, 1, seq_len, seq_len]
# 计算注意力
attn_output = scaled_dot_product_attention(q, k, v, mask)
# 合并多头
attn_output = attn_output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
return self.w_o(attn_output)
4. 实战:构建简易Transformer层
现在我们将位置编码和Mask机制整合到一个简易的Transformer层中。
4.1 完整层实现
class TransformerLayer(torch.nn.Module):
def __init__(self, d_model: int, n_heads: int, ff_dim: int, dropout: float = 0.1):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, n_heads)
self.ffn = torch.nn.Sequential(
torch.nn.Linear(d_model, ff_dim),
torch.nn.ReLU(),
torch.nn.Linear(ff_dim, d_model)
)
self.norm1 = torch.nn.LayerNorm(d_model)
self.norm2 = torch.nn.LayerNorm(d_model)
self.dropout = torch.nn.Dropout(dropout)
def forward(self, x: torch.Tensor, mask: torch.Tensor = None) -> torch.Tensor:
# 自注意力
attn_output = self.self_attn(x, x, x, mask)
x = x + self.dropout(attn_output)
x = self.norm1(x)
# 前馈网络
ffn_output = self.ffn(x)
x = x + self.dropout(ffn_output)
x = self.norm2(x)
return x
4.2 使用示例
# 参数设置
d_model = 512
n_heads = 8
ff_dim = 2048
seq_len = 100
batch_size = 32
vocab_size = 10000
# 初始化组件
embedding = torch.nn.Embedding(vocab_size, d_model)
pos_encoder = PositionalEncoding(d_model)
transformer_layer = TransformerLayer(d_model, n_heads, ff_dim)
# 模拟输入
input_seq = torch.randint(0, vocab_size, (batch_size, seq_len))
pad_mask = create_pad_mask(input_seq, pad_idx=0) # 假设0是pad索引
seq_mask = create_seq_mask(seq_len)
# 前向传播
x = embedding(input_seq)
x = pos_encoder(x)
output = transformer_layer(x, mask=combine_masks(pad_mask, seq_mask))
4.3 调试技巧
在实现Transformer组件时,有几个调试技巧很有用:
-
形状检查 :在每个关键步骤打印张量形状
print(f"Shape after embedding: {x.shape}") -
Mask验证 :可视化Mask确保其正确性
plt.imshow(pad_mask[0, 0, 0].cpu().numpy()) plt.title('Pad Mask') plt.show() -
梯度检查 :使用PyTorch的autograd.gradcheck验证反向传播
from torch.autograd import gradcheck test_input = torch.randn(2, 10, d_model, dtype=torch.double, requires_grad=True) test = gradcheck(TransformerLayer(d_model, n_heads, ff_dim), test_input) print("Gradient check passed:", test)
通过本文的代码实现,你应该对Transformer的位置编码和Mask机制有了更直观的理解。这些核心组件不仅是Transformer成功的关键,也是理解现代NLP模型的基础。
更多推荐




所有评论(0)